diff --git a/google/cloud/aiplatform/v1/schema/predict/instance_v1/__init__.py b/google/cloud/aiplatform/v1/schema/predict/instance_v1/__init__.py index 36b1a19a08..b9ba40a62f 100644 --- a/google/cloud/aiplatform/v1/schema/predict/instance_v1/__init__.py +++ b/google/cloud/aiplatform/v1/schema/predict/instance_v1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -53,89 +50,6 @@ from .types.video_classification import VideoClassificationPredictionInstance from .types.video_object_tracking import VideoObjectTrackingPredictionInstance -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1.schema.predict.instance_v1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1.schema.predict.instance_v1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = "google.cloud.aiplatform.v1.schema.predict.instance_v1" - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "ImageClassificationPredictionInstance", "ImageObjectDetectionPredictionInstance", @@ -147,3 +61,8 @@ def _get_version(dependency_name): "VideoClassificationPredictionInstance", "VideoObjectTrackingPredictionInstance", ) + +api_core.check_python_version("google.cloud.aiplatform.v1.schema.predict.instance_v1") +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1.schema.predict.instance_v1" +) diff --git a/google/cloud/aiplatform/v1/schema/predict/instance_v1/_compat.py b/google/cloud/aiplatform/v1/schema/predict/instance_v1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1/schema/predict/instance_v1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform/v1/schema/predict/params_v1/__init__.py b/google/cloud/aiplatform/v1/schema/predict/params_v1/__init__.py index 058520b481..0b3881b83a 100644 --- a/google/cloud/aiplatform/v1/schema/predict/params_v1/__init__.py +++ b/google/cloud/aiplatform/v1/schema/predict/params_v1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -47,89 +44,6 @@ from .types.video_classification import VideoClassificationPredictionParams from .types.video_object_tracking import VideoObjectTrackingPredictionParams -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1.schema.predict.params_v1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1.schema.predict.params_v1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = "google.cloud.aiplatform.v1.schema.predict.params_v1" - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "ImageClassificationPredictionParams", "ImageObjectDetectionPredictionParams", @@ -138,3 +52,8 @@ def _get_version(dependency_name): "VideoClassificationPredictionParams", "VideoObjectTrackingPredictionParams", ) + +api_core.check_python_version("google.cloud.aiplatform.v1.schema.predict.params_v1") +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1.schema.predict.params_v1" +) diff --git a/google/cloud/aiplatform/v1/schema/predict/params_v1/_compat.py b/google/cloud/aiplatform/v1/schema/predict/params_v1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1/schema/predict/params_v1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform/v1/schema/predict/prediction_v1/__init__.py b/google/cloud/aiplatform/v1/schema/predict/prediction_v1/__init__.py index 973b0a123e..1ed366ba75 100644 --- a/google/cloud/aiplatform/v1/schema/predict/prediction_v1/__init__.py +++ b/google/cloud/aiplatform/v1/schema/predict/prediction_v1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -55,89 +52,6 @@ from .types.video_classification import VideoClassificationPredictionResult from .types.video_object_tracking import VideoObjectTrackingPredictionResult -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1.schema.predict.prediction_v1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1.schema.predict.prediction_v1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = "google.cloud.aiplatform.v1.schema.predict.prediction_v1" - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "ClassificationPredictionResult", "ImageObjectDetectionPredictionResult", @@ -150,3 +64,8 @@ def _get_version(dependency_name): "VideoClassificationPredictionResult", "VideoObjectTrackingPredictionResult", ) + +api_core.check_python_version("google.cloud.aiplatform.v1.schema.predict.prediction_v1") +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1.schema.predict.prediction_v1" +) diff --git a/google/cloud/aiplatform/v1/schema/predict/prediction_v1/_compat.py b/google/cloud/aiplatform/v1/schema/predict/prediction_v1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1/schema/predict/prediction_v1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform/v1/schema/trainingjob/definition_v1/__init__.py b/google/cloud/aiplatform/v1/schema/trainingjob/definition_v1/__init__.py index 9bb4a8f193..0f522bf874 100644 --- a/google/cloud/aiplatform/v1/schema/trainingjob/definition_v1/__init__.py +++ b/google/cloud/aiplatform/v1/schema/trainingjob/definition_v1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -71,89 +68,6 @@ from .types.automl_video_object_tracking import AutoMlVideoObjectTrackingInputs from .types.export_evaluated_data_items_config import ExportEvaluatedDataItemsConfig -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1.schema.trainingjob.definition_v1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1.schema.trainingjob.definition_v1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = "google.cloud.aiplatform.v1.schema.trainingjob.definition_v1" - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "AutoMlImageClassification", "AutoMlImageClassificationInputs", @@ -181,3 +95,10 @@ def _get_version(dependency_name): "AutoMlVideoObjectTrackingInputs", "ExportEvaluatedDataItemsConfig", ) + +api_core.check_python_version( + "google.cloud.aiplatform.v1.schema.trainingjob.definition_v1" +) +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1.schema.trainingjob.definition_v1" +) diff --git a/google/cloud/aiplatform/v1/schema/trainingjob/definition_v1/_compat.py b/google/cloud/aiplatform/v1/schema/trainingjob/definition_v1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1/schema/trainingjob/definition_v1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform/v1beta1/schema/predict/instance_v1beta1/__init__.py b/google/cloud/aiplatform/v1beta1/schema/predict/instance_v1beta1/__init__.py index 3e184cd5eb..2641b134f3 100644 --- a/google/cloud/aiplatform/v1beta1/schema/predict/instance_v1beta1/__init__.py +++ b/google/cloud/aiplatform/v1beta1/schema/predict/instance_v1beta1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -53,91 +50,6 @@ from .types.video_classification import VideoClassificationPredictionInstance from .types.video_object_tracking import VideoObjectTrackingPredictionInstance -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1beta1.schema.predict.instance_v1beta1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1beta1.schema.predict.instance_v1beta1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = ( - "google.cloud.aiplatform.v1beta1.schema.predict.instance_v1beta1" - ) - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "ImageClassificationPredictionInstance", "ImageObjectDetectionPredictionInstance", @@ -149,3 +61,10 @@ def _get_version(dependency_name): "VideoClassificationPredictionInstance", "VideoObjectTrackingPredictionInstance", ) + +api_core.check_python_version( + "google.cloud.aiplatform.v1beta1.schema.predict.instance_v1beta1" +) +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1beta1.schema.predict.instance_v1beta1" +) diff --git a/google/cloud/aiplatform/v1beta1/schema/predict/instance_v1beta1/_compat.py b/google/cloud/aiplatform/v1beta1/schema/predict/instance_v1beta1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1beta1/schema/predict/instance_v1beta1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform/v1beta1/schema/predict/params_v1beta1/__init__.py b/google/cloud/aiplatform/v1beta1/schema/predict/params_v1beta1/__init__.py index 82e9efab3d..73ca8f08a9 100644 --- a/google/cloud/aiplatform/v1beta1/schema/predict/params_v1beta1/__init__.py +++ b/google/cloud/aiplatform/v1beta1/schema/predict/params_v1beta1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -47,89 +44,6 @@ from .types.video_classification import VideoClassificationPredictionParams from .types.video_object_tracking import VideoObjectTrackingPredictionParams -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1beta1.schema.predict.params_v1beta1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1beta1.schema.predict.params_v1beta1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = "google.cloud.aiplatform.v1beta1.schema.predict.params_v1beta1" - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "ImageClassificationPredictionParams", "ImageObjectDetectionPredictionParams", @@ -138,3 +52,10 @@ def _get_version(dependency_name): "VideoClassificationPredictionParams", "VideoObjectTrackingPredictionParams", ) + +api_core.check_python_version( + "google.cloud.aiplatform.v1beta1.schema.predict.params_v1beta1" +) +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1beta1.schema.predict.params_v1beta1" +) diff --git a/google/cloud/aiplatform/v1beta1/schema/predict/params_v1beta1/_compat.py b/google/cloud/aiplatform/v1beta1/schema/predict/params_v1beta1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1beta1/schema/predict/params_v1beta1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform/v1beta1/schema/predict/prediction_v1beta1/__init__.py b/google/cloud/aiplatform/v1beta1/schema/predict/prediction_v1beta1/__init__.py index e775eb9c68..0a3edf1db9 100644 --- a/google/cloud/aiplatform/v1beta1/schema/predict/prediction_v1beta1/__init__.py +++ b/google/cloud/aiplatform/v1beta1/schema/predict/prediction_v1beta1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -57,91 +54,6 @@ from .types.video_classification import VideoClassificationPredictionResult from .types.video_object_tracking import VideoObjectTrackingPredictionResult -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1beta1.schema.predict.prediction_v1beta1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1beta1.schema.predict.prediction_v1beta1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = ( - "google.cloud.aiplatform.v1beta1.schema.predict.prediction_v1beta1" - ) - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "ClassificationPredictionResult", "ImageObjectDetectionPredictionResult", @@ -155,3 +67,10 @@ def _get_version(dependency_name): "VideoClassificationPredictionResult", "VideoObjectTrackingPredictionResult", ) + +api_core.check_python_version( + "google.cloud.aiplatform.v1beta1.schema.predict.prediction_v1beta1" +) +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1beta1.schema.predict.prediction_v1beta1" +) diff --git a/google/cloud/aiplatform/v1beta1/schema/predict/prediction_v1beta1/_compat.py b/google/cloud/aiplatform/v1beta1/schema/predict/prediction_v1beta1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1beta1/schema/predict/prediction_v1beta1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform/v1beta1/schema/trainingjob/definition_v1beta1/__init__.py b/google/cloud/aiplatform/v1beta1/schema/trainingjob/definition_v1beta1/__init__.py index 3b17816244..b5250a01e7 100644 --- a/google/cloud/aiplatform/v1beta1/schema/trainingjob/definition_v1beta1/__init__.py +++ b/google/cloud/aiplatform/v1beta1/schema/trainingjob/definition_v1beta1/__init__.py @@ -18,12 +18,9 @@ ) import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -75,91 +72,6 @@ from .types.automl_video_object_tracking import AutoMlVideoObjectTrackingInputs from .types.export_evaluated_data_items_config import ExportEvaluatedDataItemsConfig -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform.v1beta1.schema.trainingjob.definition_v1beta1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform.v1beta1.schema.trainingjob.definition_v1beta1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = ( - "google.cloud.aiplatform.v1beta1.schema.trainingjob.definition_v1beta1" - ) - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "AutoMlForecasting", "AutoMlForecastingInputs", @@ -190,3 +102,10 @@ def _get_version(dependency_name): "AutoMlVideoObjectTrackingInputs", "ExportEvaluatedDataItemsConfig", ) + +api_core.check_python_version( + "google.cloud.aiplatform.v1beta1.schema.trainingjob.definition_v1beta1" +) +api_core.check_dependency_versions( + "google.cloud.aiplatform.v1beta1.schema.trainingjob.definition_v1beta1" +) diff --git a/google/cloud/aiplatform/v1beta1/schema/trainingjob/definition_v1beta1/_compat.py b/google/cloud/aiplatform/v1beta1/schema/trainingjob/definition_v1beta1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform/v1beta1/schema/trainingjob/definition_v1beta1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform_v1/__init__.py b/google/cloud/aiplatform_v1/__init__.py index 90f2e37f7c..73a8be2c23 100644 --- a/google/cloud/aiplatform_v1/__init__.py +++ b/google/cloud/aiplatform_v1/__init__.py @@ -16,12 +16,9 @@ from google.cloud.aiplatform_v1 import gapic_version as package_version import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -1306,89 +1303,6 @@ from .types.vizier_service import SuggestTrialsRequest from .types.vizier_service import SuggestTrialsResponse -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform_v1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform_v1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = "google.cloud.aiplatform_v1" - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "DataFoundryServiceAsyncClient", "DatasetServiceAsyncClient", @@ -2467,3 +2381,6 @@ def _get_version(dependency_name): "WriteTensorboardRunDataResponse", "XraiAttribution", ) + +api_core.check_python_version("google.cloud.aiplatform_v1") +api_core.check_dependency_versions("google.cloud.aiplatform_v1") diff --git a/google/cloud/aiplatform_v1/_compat.py b/google/cloud/aiplatform_v1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform_v1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform_v1/services/data_foundry_service/client.py b/google/cloud/aiplatform_v1/services/data_foundry_service/client.py index 3850036a0f..286e5cacf5 100644 --- a/google/cloud/aiplatform_v1/services/data_foundry_service/client.py +++ b/google/cloud/aiplatform_v1/services/data_foundry_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -129,77 +136,13 @@ class DataFoundryServiceClient(metaclass=DataFoundryServiceClientMeta): evaluation. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -365,7 +308,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = DataFoundryServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -386,35 +329,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = DataFoundryServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -434,65 +354,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = DataFoundryServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = DataFoundryServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = DataFoundryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = DataFoundryServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -626,13 +487,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - DataFoundryServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = DataFoundryServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = DataFoundryServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=DataFoundryServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -668,14 +531,14 @@ def __init__( self._transport = cast(DataFoundryServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or DataFoundryServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=DataFoundryServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=DataFoundryServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=DataFoundryServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest.py b/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest.py index 3aa45dc3ca..afdc1fc2dd 100644 --- a/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -548,21 +549,18 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_http_options() ) - request, metadata = self._interceptor.pre_generate_synthetic_data( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_transcoded_request( - http_options, request - ) - - body = _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData, + "_BaseGenerateSyntheticData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -721,15 +719,16 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -863,15 +862,16 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1005,15 +1005,16 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1148,19 +1149,16 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1295,17 +1293,18 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1438,17 +1437,18 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1554,17 +1554,18 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1673,15 +1674,16 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1815,15 +1817,16 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1957,15 +1960,16 @@ def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_asyncio.py index 01783a3956..7af7d6a075 100644 --- a/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -603,21 +605,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_http_options() ) - request, metadata = await self._interceptor.pre_generate_synthetic_data( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_transcoded_request( - http_options, request - ) - - body = _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData, + "_BaseGenerateSyntheticData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -780,17 +779,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -932,17 +932,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1084,17 +1085,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1237,21 +1239,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1394,17 +1393,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1541,17 +1541,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1663,17 +1664,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1788,17 +1790,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1940,17 +1943,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2090,17 +2094,18 @@ async def __call__( http_options = ( _BaseDataFoundryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseDataFoundryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDataFoundryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDataFoundryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_base.py index 83eaba7d2a..05c50b72d6 100644 --- a/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/data_foundry_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,38 +110,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = data_foundry_service.GenerateSyntheticDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDataFoundryServiceRestTransport._BaseGenerateSyntheticData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -168,17 +128,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -197,17 +146,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -278,17 +216,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -369,22 +296,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -451,17 +362,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -836,17 +736,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1249,17 +1138,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1674,17 +1552,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2095,17 +1962,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2512,16 +2368,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseDataFoundryServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/dataset_service/client.py b/google/cloud/aiplatform_v1/services/dataset_service/client.py index 1a20896c0b..3bdedc743f 100644 --- a/google/cloud/aiplatform_v1/services/dataset_service/client.py +++ b/google/cloud/aiplatform_v1/services/dataset_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -146,77 +153,13 @@ class DatasetServiceClient(metaclass=DatasetServiceClientMeta): resources. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -526,7 +469,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = DatasetServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -547,35 +490,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = DatasetServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -595,65 +515,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = DatasetServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = DatasetServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -783,13 +644,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - DatasetServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = DatasetServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = DatasetServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=DatasetServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -825,14 +688,14 @@ def __init__( self._transport = cast(DatasetServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or DatasetServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=DatasetServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=DatasetServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/dataset_service/transports/rest.py b/google/cloud/aiplatform_v1/services/dataset_service/transports/rest.py index 004546c724..d6c6204f5a 100644 --- a/google/cloud/aiplatform_v1/services/dataset_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/dataset_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3572,19 +3573,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDataset._get_http_options() ) - request, metadata = self._interceptor.pre_create_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDataset, + "_BaseCreateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3726,21 +3724,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_create_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion, + "_BaseCreateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3880,15 +3875,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_http_options() ) - request, metadata = self._interceptor.pre_delete_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDataset, + "_BaseDeleteDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4028,17 +4024,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_delete_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion, + "_BaseDeleteDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4177,17 +4174,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_http_options() ) - request, metadata = self._interceptor.pre_delete_saved_query( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery, + "_BaseDeleteSavedQuery__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4327,23 +4325,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseExportData._get_http_options() ) - request, metadata = self._interceptor.pre_export_data(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseExportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseExportData, + "_BaseExportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4482,17 +4473,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_http_options() ) - request, metadata = self._interceptor.pre_get_annotation_spec( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec, + "_BaseGetAnnotationSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4632,17 +4624,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDataset._get_http_options() ) - request, metadata = self._interceptor.pre_get_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseGetDataset._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDataset, + "_BaseGetDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4780,17 +4771,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_get_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDatasetVersion, + "_BaseGetDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4932,23 +4924,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseImportData._get_http_options() ) - request, metadata = self._interceptor.pre_import_data(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseImportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseImportData, + "_BaseImportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5087,17 +5072,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListAnnotations._get_http_options() ) - request, metadata = self._interceptor.pre_list_annotations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListAnnotations, + "_BaseListAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5239,15 +5225,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDataItems._get_http_options() ) - request, metadata = self._interceptor.pre_list_data_items(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDataItems, + "_BaseListDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5389,15 +5376,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasets._get_http_options() ) - request, metadata = self._interceptor.pre_list_datasets(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasets._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasets._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasets, + "_BaseListDatasets__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5540,17 +5528,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_http_options() ) - request, metadata = self._interceptor.pre_list_dataset_versions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasetVersions, + "_BaseListDatasetVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5692,17 +5681,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_http_options() ) - request, metadata = self._interceptor.pre_list_saved_queries( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListSavedQueries, + "_BaseListSavedQueries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5846,17 +5836,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_restore_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion, + "_BaseRestoreDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5994,17 +5985,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_http_options() ) - request, metadata = self._interceptor.pre_search_data_items( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSearchDataItems, + "_BaseSearchDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6147,19 +6139,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_http_options() ) - request, metadata = self._interceptor.pre_update_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDataset, + "_BaseUpdateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6300,21 +6289,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_update_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion, + "_BaseUpdateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6638,15 +6624,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6779,15 +6766,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6920,15 +6908,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7062,19 +7051,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7208,17 +7194,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7348,17 +7335,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7463,17 +7451,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7581,15 +7570,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7722,15 +7712,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7863,15 +7854,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_asyncio.py index ffe43e1c87..e3b8342379 100644 --- a/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1728,21 +1730,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_create_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDataset, + "_BaseCreateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1894,21 +1893,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_create_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion, + "_BaseCreateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2056,17 +2052,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDataset, + "_BaseDeleteDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2216,17 +2213,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion, + "_BaseDeleteDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2374,17 +2372,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_saved_query( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery, + "_BaseDeleteSavedQuery__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2534,25 +2533,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseExportData._get_http_options() ) - request, metadata = await self._interceptor.pre_export_data( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseExportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseExportData, + "_BaseExportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2700,17 +2692,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_http_options() ) - request, metadata = await self._interceptor.pre_get_annotation_spec( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec, + "_BaseGetAnnotationSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2858,19 +2851,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_get_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseGetDataset._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDataset, + "_BaseGetDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3015,17 +3007,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_get_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDatasetVersion, + "_BaseGetDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3175,25 +3168,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseImportData._get_http_options() ) - request, metadata = await self._interceptor.pre_import_data( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseImportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseImportData, + "_BaseImportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3341,17 +3327,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListAnnotations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_annotations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListAnnotations, + "_BaseListAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3501,17 +3488,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDataItems._get_http_options() ) - request, metadata = await self._interceptor.pre_list_data_items( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDataItems, + "_BaseListDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3661,17 +3649,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasets._get_http_options() ) - request, metadata = await self._interceptor.pre_list_datasets( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasets._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasets._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasets, + "_BaseListDatasets__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3822,17 +3811,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_dataset_versions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasetVersions, + "_BaseListDatasetVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3981,17 +3971,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_http_options() ) - request, metadata = await self._interceptor.pre_list_saved_queries( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListSavedQueries, + "_BaseListSavedQueries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4143,17 +4134,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_restore_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion, + "_BaseRestoreDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4302,17 +4294,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_http_options() ) - request, metadata = await self._interceptor.pre_search_data_items( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSearchDataItems, + "_BaseSearchDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4463,21 +4456,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_update_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDataset, + "_BaseUpdateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4626,21 +4616,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_update_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion, + "_BaseUpdateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6925,17 +6912,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7076,17 +7064,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7227,17 +7216,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7379,21 +7369,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7536,17 +7523,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7683,17 +7671,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7807,17 +7796,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7933,17 +7923,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8085,17 +8076,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8236,17 +8228,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_base.py index 034ecca6e8..7047bf5792 100644 --- a/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/dataset_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -128,52 +120,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.CreateDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseCreateDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -190,52 +142,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.CreateDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -250,43 +162,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.DeleteDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -301,43 +182,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.DeleteDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSavedQuery: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -348,43 +198,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.DeleteSavedQueryRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -396,52 +215,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ExportDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseExportData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetAnnotationSpec: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -452,43 +231,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.GetAnnotationSpecRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -503,43 +251,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.GetDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseGetDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -554,43 +271,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.GetDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -602,52 +288,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ImportDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseImportData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListAnnotations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -658,43 +304,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListAnnotationsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListAnnotations._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDataItems: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -705,43 +320,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListDataItemsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListDataItems._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDatasets: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -756,43 +340,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListDatasetsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListDatasets._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDatasetVersions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -807,43 +360,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListDatasetVersionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSavedQueries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -854,43 +376,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListSavedQueriesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRestoreDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -905,43 +396,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.RestoreDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchDataItems: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -952,29 +412,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.SearchDataItemsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -983,14 +420,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1007,38 +436,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.UpdateDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1047,14 +444,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1071,38 +460,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.UpdateDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1121,17 +478,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1150,17 +496,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1231,17 +566,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1322,22 +646,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1404,17 +712,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1789,17 +1086,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2202,17 +1488,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2627,17 +1902,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3048,17 +2312,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3465,16 +2718,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseDatasetServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/client.py b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/client.py index 53297e462d..fd3772beb7 100644 --- a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/client.py +++ b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -147,77 +154,13 @@ class DeploymentResourcePoolServiceClient( ): """A service that manages the DeploymentResourcePool resource.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -471,9 +414,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - DeploymentResourcePoolServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -494,37 +435,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - DeploymentResourcePoolServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -544,67 +460,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -738,17 +593,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - DeploymentResourcePoolServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( DeploymentResourcePoolServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - DeploymentResourcePoolServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -786,14 +641,14 @@ def __init__( self._transport = cast(DeploymentResourcePoolServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or DeploymentResourcePoolServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest.py b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest.py index 9496f56da8..95302e941f 100644 --- a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2844,21 +2845,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_create_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool, + "_BaseCreateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3004,17 +3002,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_delete_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool, + "_BaseDeleteDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3160,17 +3159,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_get_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool, + "_BaseGetDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3318,17 +3318,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_http_options() ) - request, metadata = self._interceptor.pre_list_deployment_resource_pools( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools, + "_BaseListDeploymentResourcePools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3481,17 +3482,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_http_options() ) - request, metadata = self._interceptor.pre_query_deployed_models( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels, + "_BaseQueryDeployedModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3641,21 +3643,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_update_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool, + "_BaseUpdateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3863,15 +3862,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4007,15 +4007,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4151,15 +4152,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4296,19 +4298,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4445,17 +4444,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4586,17 +4586,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4702,17 +4703,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4821,15 +4823,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4965,15 +4968,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5107,15 +5111,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_asyncio.py index 1cfdc0f904..87a59cdc42 100644 --- a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -933,23 +935,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool, + "_BaseCreateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1103,19 +1102,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool, + "_BaseDeleteDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1269,19 +1269,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool, + "_BaseGetDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1437,19 +1438,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_deployment_resource_pools( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools, + "_BaseListDeploymentResourcePools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1608,17 +1610,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_http_options() ) - request, metadata = await self._interceptor.pre_query_deployed_models( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels, + "_BaseQueryDeployedModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1774,23 +1777,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool, + "_BaseUpdateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3987,17 +3987,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4137,17 +4138,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4287,17 +4289,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4438,21 +4441,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4595,17 +4595,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4744,17 +4745,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4868,17 +4870,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4993,17 +4996,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5145,17 +5149,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5295,17 +5300,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_base.py index b776f53b48..1fa847f7d0 100644 --- a/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/deployment_resource_pool_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,56 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.CreateDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDeploymentResourcePool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -181,47 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.DeleteDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDeploymentResourcePool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -232,47 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.GetDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDeploymentResourcePools: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -283,47 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.ListDeploymentResourcePoolsRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryDeployedModels: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -334,31 +177,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = deployment_resource_pool_service.QueryDeployedModelsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateDeploymentResourcePool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -367,14 +185,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -386,42 +196,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.UpdateDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -440,17 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -469,17 +232,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -550,17 +302,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -641,22 +382,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -723,17 +448,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1108,17 +822,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1521,17 +1224,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1946,17 +1638,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2367,17 +2048,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2784,16 +2454,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseDeploymentResourcePoolServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/endpoint_service/client.py b/google/cloud/aiplatform_v1/services/endpoint_service/client.py index cc9d873a64..6981e885f9 100644 --- a/google/cloud/aiplatform_v1/services/endpoint_service/client.py +++ b/google/cloud/aiplatform_v1/services/endpoint_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class EndpointServiceClient(metaclass=EndpointServiceClientMeta): """A service for managing Vertex AI's Endpoints.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -503,7 +446,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = EndpointServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -524,35 +467,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = EndpointServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -572,65 +492,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = EndpointServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = EndpointServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -762,13 +623,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - EndpointServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = EndpointServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = EndpointServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=EndpointServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -804,14 +667,14 @@ def __init__( self._transport = cast(EndpointServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or EndpointServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=EndpointServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=EndpointServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest.py b/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest.py index 7bd8243142..c71135c77a 100644 --- a/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2991,19 +2992,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_create_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCreateEndpoint, + "_BaseCreateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3143,15 +3141,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_delete_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteEndpoint, + "_BaseDeleteEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3291,19 +3290,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeployModel._get_http_options() ) - request, metadata = self._interceptor.pre_deploy_model(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseDeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeployModel, + "_BaseDeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3443,15 +3439,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_get_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetEndpoint, + "_BaseGetEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3591,15 +3588,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListEndpoints._get_http_options() ) - request, metadata = self._interceptor.pre_list_endpoints(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListEndpoints, + "_BaseListEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3744,21 +3742,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_http_options() ) - request, metadata = self._interceptor.pre_mutate_deployed_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseMutateDeployedModel, + "_BaseMutateDeployedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3899,19 +3894,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUndeployModel._get_http_options() ) - request, metadata = self._interceptor.pre_undeploy_model(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUndeployModel, + "_BaseUndeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4052,19 +4044,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_update_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpoint, + "_BaseUpdateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4209,21 +4198,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_http_options() ) - request, metadata = self._interceptor.pre_update_endpoint_long_running( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning, + "_BaseUpdateEndpointLongRunning__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4442,15 +4428,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4583,15 +4570,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4724,15 +4712,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4866,19 +4855,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5013,17 +4999,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5153,17 +5140,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5268,17 +5256,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5386,15 +5375,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5527,15 +5517,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5668,15 +5659,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_asyncio.py index 5303cac03b..2f735d23f1 100644 --- a/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1096,21 +1098,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_create_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCreateEndpoint, + "_BaseCreateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1261,17 +1260,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteEndpoint, + "_BaseDeleteEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1421,21 +1421,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeployModel._get_http_options() ) - request, metadata = await self._interceptor.pre_deploy_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseDeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeployModel, + "_BaseDeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1585,17 +1582,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_get_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetEndpoint, + "_BaseGetEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1744,17 +1742,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListEndpoints._get_http_options() ) - request, metadata = await self._interceptor.pre_list_endpoints( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListEndpoints, + "_BaseListEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1907,21 +1906,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_http_options() ) - request, metadata = await self._interceptor.pre_mutate_deployed_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseMutateDeployedModel, + "_BaseMutateDeployedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2071,21 +2067,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUndeployModel._get_http_options() ) - request, metadata = await self._interceptor.pre_undeploy_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUndeployModel, + "_BaseUndeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2237,21 +2230,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_update_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpoint, + "_BaseUpdateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2404,23 +2394,20 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_endpoint_long_running( request, metadata ) ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning, + "_BaseUpdateEndpointLongRunning__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4622,17 +4609,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4774,17 +4762,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4926,17 +4915,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5079,21 +5069,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5236,17 +5223,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5383,17 +5371,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5507,17 +5496,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5634,17 +5624,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5786,17 +5777,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5938,17 +5930,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_base.py index 810a2073a0..7aa341482f 100644 --- a/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/endpoint_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,52 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.CreateEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -176,43 +128,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.DeleteEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeployModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,52 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.DeployModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseDeployModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -280,43 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.GetEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListEndpoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -327,43 +177,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.ListEndpointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseListEndpoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseMutateDeployedModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -375,52 +194,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.MutateDeployedModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUndeployModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -432,38 +211,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.UndeployModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseUndeployModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -472,14 +219,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -491,52 +230,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.UpdateEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateEndpointLongRunning: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -548,38 +247,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.UpdateEndpointLongRunningRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -598,17 +265,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -627,17 +283,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -708,17 +353,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -799,22 +433,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -881,17 +499,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1266,17 +873,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1679,17 +1275,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2104,17 +1689,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2525,17 +2099,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2942,16 +2505,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseEndpointServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/evaluation_service/client.py b/google/cloud/aiplatform_v1/services/evaluation_service/client.py index 64f85cc732..df5a6caa3b 100644 --- a/google/cloud/aiplatform_v1/services/evaluation_service/client.py +++ b/google/cloud/aiplatform_v1/services/evaluation_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -127,77 +134,13 @@ def get_transport_class( class EvaluationServiceClient(metaclass=EvaluationServiceClientMeta): """Vertex AI Online Evaluation Service.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -363,7 +306,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = EvaluationServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -384,35 +327,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = EvaluationServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -432,65 +352,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = EvaluationServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = EvaluationServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -624,13 +485,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - EvaluationServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = EvaluationServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = EvaluationServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=EvaluationServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -666,14 +529,14 @@ def __init__( self._transport = cast(EvaluationServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or EvaluationServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=EvaluationServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest.py b/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest.py index dbd75c81e5..6bd9f23d9c 100644 --- a/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -547,21 +548,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_http_options() ) - request, metadata = self._interceptor.pre_evaluate_instances( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseEvaluateInstances, + "_BaseEvaluateInstances__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -715,15 +713,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -857,15 +856,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -998,15 +998,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1140,19 +1141,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1287,17 +1285,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1428,17 +1427,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1544,17 +1544,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1662,15 +1663,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1804,15 +1806,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1946,15 +1949,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_asyncio.py index f906239886..0407b64a26 100644 --- a/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -602,21 +604,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_http_options() ) - request, metadata = await self._interceptor.pre_evaluate_instances( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseEvaluateInstances, + "_BaseEvaluateInstances__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -775,17 +774,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -927,17 +927,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1079,17 +1080,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1232,21 +1234,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1389,17 +1388,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1536,17 +1536,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1658,17 +1659,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1783,17 +1785,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1935,17 +1938,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2087,17 +2091,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_base.py index 23b60679f7..60ab68a417 100644 --- a/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/evaluation_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,38 +110,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = evaluation_service.EvaluateInstancesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -168,17 +128,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -197,17 +146,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -278,17 +216,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -369,22 +296,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -451,17 +362,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -836,17 +736,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1249,17 +1138,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1674,17 +1552,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2095,17 +1962,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2512,16 +2368,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseEvaluationServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/client.py b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/client.py index ca05170688..4eba665513 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/client.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -159,77 +166,13 @@ class FeatureOnlineStoreAdminServiceClient( FeatureOnlineStore. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -465,9 +408,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -488,37 +429,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -538,67 +454,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -732,17 +587,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeatureOnlineStoreAdminServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( FeatureOnlineStoreAdminServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - FeatureOnlineStoreAdminServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -780,14 +635,14 @@ def __init__( self._transport = cast(FeatureOnlineStoreAdminServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeatureOnlineStoreAdminServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest.py b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest.py index ed85a7e3ee..392ca3c89d 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3253,21 +3254,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore, + "_BaseCreateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3409,21 +3407,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView, + "_BaseCreateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3567,17 +3562,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore, + "_BaseDeleteFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3717,17 +3713,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView, + "_BaseDeleteFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3871,17 +3868,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore, + "_BaseGetFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4025,17 +4023,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView, + "_BaseGetFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4180,17 +4179,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_view_sync( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync, + "_BaseGetFeatureViewSync__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4336,17 +4336,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_online_stores( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores, + "_BaseListFeatureOnlineStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4493,17 +4494,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_views( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews, + "_BaseListFeatureViews__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4650,17 +4652,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_view_syncs( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs, + "_BaseListFeatureViewSyncs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4806,21 +4809,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_sync_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView, + "_BaseSyncFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4971,21 +4971,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore, + "_BaseUpdateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5127,21 +5124,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView, + "_BaseUpdateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5424,15 +5418,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5568,15 +5563,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5710,15 +5706,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5855,19 +5852,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6006,17 +6000,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6147,17 +6142,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6263,17 +6259,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6382,15 +6379,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6526,15 +6524,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6668,15 +6667,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_asyncio.py index 7eecb4c8e9..627753feee 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1379,21 +1381,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore, + "_BaseCreateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1547,21 +1546,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView, + "_BaseCreateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1713,17 +1709,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore, + "_BaseDeleteFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1875,17 +1872,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView, + "_BaseDeleteFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2037,17 +2035,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore, + "_BaseGetFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2201,17 +2200,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView, + "_BaseGetFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2362,17 +2362,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_view_sync( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync, + "_BaseGetFeatureViewSync__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2524,17 +2525,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_online_stores( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores, + "_BaseListFeatureOnlineStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2691,17 +2693,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_views( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews, + "_BaseListFeatureViews__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2854,17 +2857,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_view_syncs( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs, + "_BaseListFeatureViewSyncs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3020,21 +3024,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_sync_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView, + "_BaseSyncFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3191,21 +3192,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore, + "_BaseUpdateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3359,21 +3357,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView, + "_BaseUpdateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5631,17 +5626,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5783,17 +5779,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5933,17 +5930,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6084,21 +6082,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6241,17 +6236,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6390,17 +6386,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6514,17 +6511,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6639,17 +6637,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6791,17 +6790,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6943,17 +6943,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_base.py index 9446261320..461cbfa793 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_admin_service/transports/rest_base.py @@ -106,14 +106,6 @@ def __hash__(self): # pragma: NO COVER "featureOnlineStoreId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -125,42 +117,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.CreateFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -169,14 +125,6 @@ def __hash__(self): # pragma: NO COVER "featureViewId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -188,54 +136,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.CreateFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureOnlineStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -246,47 +152,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.DeleteFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -297,45 +168,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.DeleteFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureOnlineStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -346,47 +184,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.GetFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -397,45 +200,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.GetFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureViewSync: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -446,45 +216,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.GetFeatureViewSyncRequest.pb(request) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureOnlineStores: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -495,47 +232,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.ListFeatureOnlineStoresRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureViews: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -546,45 +248,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.ListFeatureViewsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureViewSyncs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -595,47 +264,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.ListFeatureViewSyncsRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSyncFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -647,54 +281,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.SyncFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeatureOnlineStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -706,56 +298,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.UpdateFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -767,40 +315,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.UpdateFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -819,17 +333,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -848,17 +351,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -929,17 +421,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1020,22 +501,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1102,17 +567,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1487,17 +941,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1900,17 +1343,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2325,17 +1757,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2746,17 +2167,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3163,16 +2573,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeatureOnlineStoreAdminServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_service/client.py b/google/cloud/aiplatform_v1/services/feature_online_store_service/client.py index 25619112b0..a37e1a2a35 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_service/client.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_service/client.py @@ -41,6 +41,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -134,77 +141,13 @@ def get_transport_class( class FeatureOnlineStoreServiceClient(metaclass=FeatureOnlineStoreServiceClientMeta): """A service for fetching feature values from the online store.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -394,7 +337,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = FeatureOnlineStoreServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -415,35 +358,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = FeatureOnlineStoreServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -463,67 +383,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -657,15 +516,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeatureOnlineStoreServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( FeatureOnlineStoreServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = FeatureOnlineStoreServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -701,14 +562,14 @@ def __init__( self._transport = cast(FeatureOnlineStoreServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeatureOnlineStoreServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest.py b/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest.py index e9c98a1835..cde3c7ebc5 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -690,21 +691,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_fetch_feature_values( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues, + "_BaseFetchFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -854,21 +852,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_http_options() ) - request, metadata = self._interceptor.pre_generate_fetch_access_token( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken, + "_BaseGenerateFetchAccessToken__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1015,21 +1010,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_http_options() ) - request, metadata = self._interceptor.pre_search_nearest_entities( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities, + "_BaseSearchNearestEntities__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1219,15 +1211,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1363,15 +1356,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1507,15 +1501,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1652,19 +1647,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1801,17 +1793,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1942,17 +1935,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2060,17 +2054,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2181,15 +2176,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2325,15 +2321,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2469,15 +2466,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_asyncio.py index e4b65b1857..15902ceb3f 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -766,21 +768,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_fetch_feature_values( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues, + "_BaseFetchFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -936,21 +935,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_http_options() ) - request, metadata = await self._interceptor.pre_generate_fetch_access_token( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken, + "_BaseGenerateFetchAccessToken__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1107,21 +1103,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_http_options() ) - request, metadata = await self._interceptor.pre_search_nearest_entities( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities, + "_BaseSearchNearestEntities__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1311,17 +1304,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1461,17 +1455,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1611,17 +1606,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1762,21 +1758,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1919,17 +1912,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2066,17 +2060,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2188,17 +2183,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2313,17 +2309,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2463,17 +2460,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2613,17 +2611,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_base.py index 6f2c5115d0..cf2586e0c6 100644 --- a/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/feature_online_store_service/transports/rest_base.py @@ -103,14 +103,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -122,40 +114,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_service.FetchFeatureValuesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGenerateFetchAccessToken: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -171,49 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_service.GenerateFetchAccessTokenRequest.pb(request) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchNearestEntities: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -225,40 +146,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_service.SearchNearestEntitiesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -277,17 +164,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -306,17 +182,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -387,17 +252,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -478,22 +332,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -560,17 +398,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -945,17 +772,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1358,17 +1174,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1783,17 +1588,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2204,17 +1998,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2621,16 +2404,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeatureOnlineStoreServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/feature_registry_service/client.py b/google/cloud/aiplatform_v1/services/feature_registry_service/client.py index c687c3c5a2..8d2ab10bd7 100644 --- a/google/cloud/aiplatform_v1/services/feature_registry_service/client.py +++ b/google/cloud/aiplatform_v1/services/feature_registry_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -141,77 +148,13 @@ class FeatureRegistryServiceClient(metaclass=FeatureRegistryServiceClientMeta): FeatureRegistry. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -425,7 +368,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = FeatureRegistryServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -446,35 +389,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = FeatureRegistryServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -494,67 +414,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeatureRegistryServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeatureRegistryServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -688,13 +547,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeatureRegistryServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = FeatureRegistryServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = FeatureRegistryServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeatureRegistryServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -730,14 +591,14 @@ def __init__( self._transport = cast(FeatureRegistryServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeatureRegistryServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeatureRegistryServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest.py b/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest.py index 9338cd12e8..08f1166f2d 100644 --- a/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3120,21 +3121,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3280,19 +3278,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3434,21 +3429,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup, + "_BaseCreateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3593,15 +3585,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3741,17 +3734,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup, + "_BaseDeleteFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3895,15 +3889,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4042,17 +4037,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup, + "_BaseGetFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4195,17 +4191,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_groups( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups, + "_BaseListFeatureGroups__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4356,15 +4353,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_list_features(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4511,19 +4509,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4665,21 +4660,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup, + "_BaseUpdateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4931,15 +4923,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5073,15 +5066,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5215,15 +5209,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5358,19 +5353,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5505,17 +5497,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5648,17 +5641,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5766,17 +5760,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5887,15 +5882,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6029,15 +6025,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6173,15 +6170,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_asyncio.py index 4f5ed2fd96..d719d5e229 100644 --- a/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1236,21 +1238,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1402,21 +1401,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1566,21 +1562,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup, + "_BaseCreateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1731,17 +1724,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1889,17 +1883,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup, + "_BaseDeleteFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2049,17 +2044,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2204,17 +2200,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup, + "_BaseGetFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2361,17 +2358,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_groups( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups, + "_BaseListFeatureGroups__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2526,17 +2524,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_list_features( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2689,21 +2688,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2853,21 +2849,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup, + "_BaseUpdateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5096,17 +5089,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5246,17 +5240,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5396,17 +5391,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5547,21 +5543,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5702,17 +5695,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5849,17 +5843,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5971,17 +5966,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6096,17 +6092,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6246,17 +6243,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6396,17 +6394,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_base.py index a25a89a60a..0753582bc5 100644 --- a/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/feature_registry_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,38 +113,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.BatchCreateFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -161,14 +121,6 @@ def __hash__(self): # pragma: NO COVER "featureId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -180,38 +132,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -220,14 +140,6 @@ def __hash__(self): # pragma: NO COVER "featureGroupId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -239,52 +151,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.CreateFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -295,43 +167,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -342,43 +183,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.DeleteFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -389,43 +199,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -436,43 +215,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.GetFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureGroups: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -483,43 +231,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.ListFeatureGroupsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatures: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -530,43 +247,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -578,52 +264,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -635,38 +281,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.UpdateFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -685,17 +299,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -714,17 +317,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -795,17 +387,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -886,22 +467,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -968,17 +533,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1353,17 +907,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1766,17 +1309,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2191,17 +1723,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2612,17 +2133,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3029,16 +2539,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeatureRegistryServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/client.py b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/client.py index 587f654f1b..fc3a5f31a1 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/client.py +++ b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/client.py @@ -40,6 +40,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -141,77 +148,13 @@ class FeaturestoreOnlineServingServiceClient( ): """A service for serving online feature values.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -401,9 +344,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - FeaturestoreOnlineServingServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -424,37 +365,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - FeaturestoreOnlineServingServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -474,65 +390,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -666,17 +523,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeaturestoreOnlineServingServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( FeaturestoreOnlineServingServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - FeaturestoreOnlineServingServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -714,14 +571,14 @@ def __init__( self._transport = cast(FeaturestoreOnlineServingServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeaturestoreOnlineServingServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest.py b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest.py index 4ffc4f15ae..b0699a5563 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -672,21 +673,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues, + "_BaseReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -837,21 +835,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_streaming_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues, + "_BaseStreamingReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -992,21 +987,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_write_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues, + "_BaseWriteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1185,15 +1177,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1327,15 +1320,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1469,15 +1463,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1612,19 +1607,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1761,17 +1753,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1902,17 +1895,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2018,17 +2012,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2137,15 +2132,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2279,15 +2275,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2421,15 +2418,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_asyncio.py index 26221d9191..4a65b59480 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -740,21 +742,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues, + "_BaseReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -910,23 +909,20 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_http_options() ) - request, metadata = ( await self._interceptor.pre_streaming_read_feature_values( request, metadata ) ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues, + "_BaseStreamingReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1073,21 +1069,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_write_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues, + "_BaseWriteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1268,17 +1261,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1420,17 +1414,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1572,17 +1567,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1725,21 +1721,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1882,17 +1875,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2031,17 +2025,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2155,17 +2150,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2282,17 +2278,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2434,17 +2431,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2586,17 +2584,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_base.py index 7a0c2e6669..e4637ad38b 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/featurestore_online_serving_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,54 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_online_service.ReadFeatureValuesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamingReadFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -179,56 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - featurestore_online_service.StreamingReadFeatureValuesRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseWriteFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -240,40 +146,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_online_service.WriteFeatureValuesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -292,17 +164,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -321,17 +182,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -402,17 +252,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -493,22 +332,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -575,17 +398,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -960,17 +772,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1373,17 +1174,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1798,17 +1588,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2219,17 +1998,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2636,16 +2404,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeaturestoreOnlineServingServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/featurestore_service/client.py b/google/cloud/aiplatform_v1/services/featurestore_service/client.py index 11ee543618..dd11ca3260 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_service/client.py +++ b/google/cloud/aiplatform_v1/services/featurestore_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -144,77 +151,13 @@ class FeaturestoreServiceClient(metaclass=FeaturestoreServiceClientMeta): Featurestore. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -452,7 +395,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = FeaturestoreServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -473,35 +416,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = FeaturestoreServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -521,65 +441,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeaturestoreServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeaturestoreServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -713,13 +574,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeaturestoreServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = FeaturestoreServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = FeaturestoreServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeaturestoreServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -755,14 +618,14 @@ def __init__( self._transport = cast(FeaturestoreServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeaturestoreServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeaturestoreServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest.py b/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest.py index 1667eab0e6..ca69fe47e1 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3696,21 +3697,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3854,21 +3852,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_batch_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues, + "_BaseBatchReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4012,21 +4007,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_create_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType, + "_BaseCreateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4170,19 +4162,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4324,21 +4313,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_create_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore, + "_BaseCreateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4481,17 +4467,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_delete_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType, + "_BaseDeleteEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4633,15 +4620,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4781,17 +4769,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_delete_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore, + "_BaseDeleteFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4934,21 +4923,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues, + "_BaseDeleteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5092,21 +5078,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_export_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues, + "_BaseExportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5252,15 +5235,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_get_entity_type(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetEntityType, + "_BaseGetEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5404,15 +5388,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5556,17 +5541,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_get_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore, + "_BaseGetFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5709,21 +5695,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_import_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues, + "_BaseImportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5865,17 +5848,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_http_options() ) - request, metadata = self._interceptor.pre_list_entity_types( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes, + "_BaseListEntityTypes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6022,15 +6006,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_list_features(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6173,17 +6158,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_http_options() ) - request, metadata = self._interceptor.pre_list_featurestores( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores, + "_BaseListFeaturestores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6328,15 +6314,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_search_features(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures, + "_BaseSearchFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6484,21 +6471,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_update_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType, + "_BaseUpdateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6644,19 +6628,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6800,21 +6781,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_update_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore, + "_BaseUpdateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7164,15 +7142,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7306,15 +7285,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7448,15 +7428,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7591,19 +7572,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7738,17 +7716,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7881,17 +7860,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7997,17 +7977,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8116,15 +8097,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8258,15 +8240,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8400,15 +8383,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_asyncio.py index c37c4085ef..b7e975b6b7 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1862,21 +1864,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2026,21 +2025,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues, + "_BaseBatchReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2192,21 +2188,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_create_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType, + "_BaseCreateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2358,21 +2351,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2522,21 +2512,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_create_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore, + "_BaseCreateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2685,17 +2672,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType, + "_BaseDeleteEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2845,17 +2833,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3003,17 +2992,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore, + "_BaseDeleteFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3162,21 +3152,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues, + "_BaseDeleteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3326,21 +3313,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_export_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues, + "_BaseExportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3492,17 +3476,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_get_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetEntityType, + "_BaseGetEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3652,17 +3637,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3814,17 +3800,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_get_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore, + "_BaseGetFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3973,21 +3960,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_import_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues, + "_BaseImportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4135,17 +4119,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_http_options() ) - request, metadata = await self._interceptor.pre_list_entity_types( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes, + "_BaseListEntityTypes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4298,17 +4283,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_list_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4459,17 +4445,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_http_options() ) - request, metadata = await self._interceptor.pre_list_featurestores( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores, + "_BaseListFeaturestores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4618,17 +4605,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_search_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures, + "_BaseSearchFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4782,21 +4770,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_update_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType, + "_BaseUpdateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4948,21 +4933,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5112,21 +5094,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_update_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore, + "_BaseUpdateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7433,17 +7412,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7585,17 +7565,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7735,17 +7716,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7888,21 +7870,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8045,17 +8024,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8192,17 +8172,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8314,17 +8295,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8439,17 +8421,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8591,17 +8574,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8741,17 +8725,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_base.py index f1f73d0ed1..555b0048b8 100644 --- a/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/featurestore_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,52 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.BatchCreateFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchReadFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -180,38 +132,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.BatchReadFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -220,14 +140,6 @@ def __hash__(self): # pragma: NO COVER "entityTypeId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -239,38 +151,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -279,14 +159,6 @@ def __hash__(self): # pragma: NO COVER "featureId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -298,38 +170,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -338,14 +178,6 @@ def __hash__(self): # pragma: NO COVER "featurestoreId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -357,52 +189,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -413,43 +205,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -460,43 +221,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -507,43 +237,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -555,52 +254,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -612,52 +271,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ExportFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -668,43 +287,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -715,43 +303,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -762,43 +319,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -810,52 +336,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ImportFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListEntityTypes: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -866,43 +352,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListEntityTypesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatures: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -913,43 +368,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeaturestores: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -960,43 +384,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListFeaturestoresRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchFeatures: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1007,43 +400,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.SearchFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1055,52 +417,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1112,52 +434,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1169,38 +451,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1219,17 +469,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1248,17 +487,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1329,17 +557,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1420,22 +637,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1502,17 +703,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1887,17 +1077,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2300,17 +1479,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2725,17 +1893,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3146,17 +2303,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3563,16 +2709,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeaturestoreServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/client.py b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/client.py index 6e88548e2a..27d7574346 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/client.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -136,77 +143,13 @@ def get_transport_class( class GenAiCacheServiceClient(metaclass=GenAiCacheServiceClientMeta): """Service for managing Vertex AI's CachedContent resource.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -416,7 +359,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = GenAiCacheServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -437,35 +380,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = GenAiCacheServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -485,65 +405,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = GenAiCacheServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = GenAiCacheServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -677,13 +538,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - GenAiCacheServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = GenAiCacheServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = GenAiCacheServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=GenAiCacheServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -719,14 +582,14 @@ def __init__( self._transport = cast(GenAiCacheServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or GenAiCacheServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=GenAiCacheServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest.py b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest.py index ee2795b2d7..4cd18c9df7 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -745,21 +746,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_create_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent, + "_BaseCreateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -899,17 +897,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_delete_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent, + "_BaseDeleteCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1019,17 +1018,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_get_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent, + "_BaseGetCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1169,17 +1169,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_http_options() ) - request, metadata = self._interceptor.pre_list_cached_contents( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListCachedContents, + "_BaseListCachedContents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1327,21 +1328,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_update_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent, + "_BaseUpdateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1537,15 +1535,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1679,15 +1678,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1820,15 +1820,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1962,19 +1963,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2109,17 +2107,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2250,17 +2249,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2366,17 +2366,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2484,15 +2485,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2626,15 +2628,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2768,15 +2771,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_asyncio.py index 2b466b0feb..f4ab32d503 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -820,21 +822,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_create_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent, + "_BaseCreateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -978,17 +977,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent, + "_BaseDeleteCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1102,17 +1102,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_get_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent, + "_BaseGetCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1258,17 +1259,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_http_options() ) - request, metadata = await self._interceptor.pre_list_cached_contents( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListCachedContents, + "_BaseListCachedContents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1422,21 +1424,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_update_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent, + "_BaseUpdateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1627,17 +1626,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1779,17 +1779,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1931,17 +1932,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2084,21 +2086,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2241,17 +2240,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2388,17 +2388,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2510,17 +2511,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2635,17 +2637,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2787,17 +2790,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2939,17 +2943,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_base.py index 2c9dbc662a..8bd39d686a 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_cache_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.CreateCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteCachedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,43 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.DeleteCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetCachedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,43 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.GetCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListCachedContents: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -271,29 +161,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.ListCachedContentsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateCachedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -302,14 +169,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -321,38 +180,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.UpdateCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -371,17 +198,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -400,17 +216,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -481,17 +286,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -572,22 +366,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -654,17 +432,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1039,17 +806,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1452,17 +1208,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1877,17 +1622,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2298,17 +2032,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2715,16 +2438,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseGenAiCacheServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/client.py b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/client.py index 63014b105a..320945aa64 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/client.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -136,77 +143,13 @@ def get_transport_class( class GenAiTuningServiceClient(metaclass=GenAiTuningServiceClientMeta): """A service for creating and managing GenAI Tuning Jobs.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -484,7 +427,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = GenAiTuningServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -505,35 +448,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = GenAiTuningServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -553,65 +473,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = GenAiTuningServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = GenAiTuningServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -745,13 +606,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - GenAiTuningServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = GenAiTuningServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = GenAiTuningServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=GenAiTuningServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -787,14 +650,14 @@ def __init__( self._transport = cast(GenAiTuningServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or GenAiTuningServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=GenAiTuningServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest.py b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest.py index 074e451e49..68ab4d3251 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2730,21 +2731,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob, + "_BaseCancelTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2853,21 +2851,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob, + "_BaseCreateTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3009,15 +3004,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_tuning_job(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob, + "_BaseGetTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3158,17 +3154,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_tuning_jobs( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs, + "_BaseListTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3313,21 +3310,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_http_options() ) - request, metadata = self._interceptor.pre_rebase_tuned_model( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel, + "_BaseRebaseTunedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3514,15 +3508,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3656,15 +3651,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3798,15 +3794,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3941,19 +3938,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4088,17 +4082,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4231,17 +4226,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4347,17 +4343,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4466,15 +4463,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4608,15 +4606,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4750,15 +4749,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_asyncio.py index 2aee740c84..2c71d79bce 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -814,21 +816,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob, + "_BaseCancelTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -943,21 +942,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob, + "_BaseCreateTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1105,17 +1101,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob, + "_BaseGetTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1264,17 +1261,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_tuning_jobs( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs, + "_BaseListTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1425,21 +1423,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_http_options() ) - request, metadata = await self._interceptor.pre_rebase_tuned_model( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel, + "_BaseRebaseTunedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3617,17 +3612,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3769,17 +3765,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3921,17 +3918,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4074,21 +4072,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4231,17 +4226,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4378,17 +4374,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4500,17 +4497,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4625,17 +4623,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4777,17 +4776,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4927,17 +4927,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_base.py index a52468d970..cf8c6e6d4d 100644 --- a/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/gen_ai_tuning_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.CancelTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -178,52 +130,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.CreateTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -234,43 +146,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.GetTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTuningJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -281,43 +162,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.ListTuningJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRebaseTunedModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -329,38 +179,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.RebaseTunedModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -379,17 +197,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -408,17 +215,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -489,17 +285,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -580,22 +365,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -662,17 +431,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1047,17 +805,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1460,17 +1207,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1885,17 +1621,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2306,17 +2031,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2723,16 +2437,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseGenAiTuningServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/index_endpoint_service/client.py b/google/cloud/aiplatform_v1/services/index_endpoint_service/client.py index 504c23e880..6a8c7d44a8 100644 --- a/google/cloud/aiplatform_v1/services/index_endpoint_service/client.py +++ b/google/cloud/aiplatform_v1/services/index_endpoint_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class IndexEndpointServiceClient(metaclass=IndexEndpointServiceClientMeta): """A service for managing Vertex AI's IndexEndpoints.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -440,7 +383,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = IndexEndpointServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -461,35 +404,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = IndexEndpointServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -509,65 +429,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = IndexEndpointServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = IndexEndpointServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -701,13 +562,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - IndexEndpointServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = IndexEndpointServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = IndexEndpointServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=IndexEndpointServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -743,14 +606,14 @@ def __init__( self._transport = cast(IndexEndpointServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or IndexEndpointServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=IndexEndpointServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest.py b/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest.py index 6a04c5fa23..eab9d1ebfc 100644 --- a/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2947,21 +2948,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_create_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint, + "_BaseCreateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3104,17 +3102,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_delete_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint, + "_BaseDeleteIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3257,19 +3256,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_http_options() ) - request, metadata = self._interceptor.pre_deploy_index(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeployIndex, + "_BaseDeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3410,17 +3406,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_get_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint, + "_BaseGetIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3563,17 +3560,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_http_options() ) - request, metadata = self._interceptor.pre_list_index_endpoints( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints, + "_BaseListIndexEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3722,21 +3720,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_http_options() ) - request, metadata = self._interceptor.pre_mutate_deployed_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex, + "_BaseMutateDeployedIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3880,19 +3875,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_http_options() ) - request, metadata = self._interceptor.pre_undeploy_index(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex, + "_BaseUndeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4034,21 +4026,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_update_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint, + "_BaseUpdateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4276,15 +4265,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4418,15 +4408,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4560,15 +4551,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4703,19 +4695,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4850,17 +4839,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4993,17 +4983,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5109,17 +5100,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5228,15 +5220,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5370,15 +5363,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5512,15 +5506,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_asyncio.py index 2f25f71d83..4d159378b4 100644 --- a/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1046,21 +1048,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_create_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint, + "_BaseCreateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1209,17 +1208,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint, + "_BaseDeleteIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1368,21 +1368,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_deploy_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeployIndex, + "_BaseDeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1533,17 +1530,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_get_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint, + "_BaseGetIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1690,17 +1688,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_http_options() ) - request, metadata = await self._interceptor.pre_list_index_endpoints( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints, + "_BaseListIndexEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1853,21 +1852,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_mutate_deployed_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex, + "_BaseMutateDeployedIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2017,21 +2013,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_undeploy_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex, + "_BaseUndeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2181,21 +2174,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_update_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint, + "_BaseUpdateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4404,17 +4394,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4556,17 +4547,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4706,17 +4698,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4857,21 +4850,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5012,17 +5002,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5159,17 +5150,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5281,17 +5273,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5406,17 +5399,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5556,17 +5550,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5706,17 +5701,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_base.py index 6a815085cc..6f1b818508 100644 --- a/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/index_endpoint_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,52 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.CreateIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteIndexEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -176,43 +128,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.DeleteIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeployIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,52 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.DeployIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetIndexEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -280,43 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.GetIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListIndexEndpoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -327,43 +177,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.ListIndexEndpointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseMutateDeployedIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -375,52 +194,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.MutateDeployedIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUndeployIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -432,38 +211,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.UndeployIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateIndexEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -472,14 +219,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -491,38 +230,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.UpdateIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -541,17 +248,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -570,17 +266,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -651,17 +336,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -742,22 +416,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -824,17 +482,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1209,17 +856,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1622,17 +1258,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2047,17 +1672,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2468,17 +2082,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2885,16 +2488,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseIndexEndpointServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/index_service/client.py b/google/cloud/aiplatform_v1/services/index_service/client.py index fcd403dcf5..4ed3ffe9bd 100644 --- a/google/cloud/aiplatform_v1/services/index_service/client.py +++ b/google/cloud/aiplatform_v1/services/index_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -139,77 +146,13 @@ class IndexServiceClient(metaclass=IndexServiceClientMeta): resources. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -419,7 +362,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = IndexServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -440,35 +383,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = IndexServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -488,65 +408,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = IndexServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = IndexServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = IndexServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -676,13 +537,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - IndexServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = IndexServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = IndexServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=IndexServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -718,11 +581,14 @@ def __init__( self._transport = cast(IndexServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or IndexServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=IndexServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=IndexServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/index_service/transports/rest.py b/google/cloud/aiplatform_v1/services/index_service/transports/rest.py index eec1b643d3..e7f32da621 100644 --- a/google/cloud/aiplatform_v1/services/index_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/index_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2877,25 +2878,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCreateIndex._get_http_options() ) - request, metadata = self._interceptor.pre_create_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCreateIndex, + "_BaseCreateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3035,19 +3027,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteIndex._get_http_options() ) - request, metadata = self._interceptor.pre_delete_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteIndex, + "_BaseDeleteIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3185,19 +3174,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIndex._get_http_options() ) - request, metadata = self._interceptor.pre_get_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIndex, + "_BaseGetIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3337,19 +3323,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListIndexes._get_http_options() ) - request, metadata = self._interceptor.pre_list_indexes(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListIndexes, + "_BaseListIndexes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3492,21 +3475,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_http_options() ) - request, metadata = self._interceptor.pre_remove_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseRemoveDatapoints, + "_BaseRemoveDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3651,25 +3631,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpdateIndex._get_http_options() ) - request, metadata = self._interceptor.pre_update_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpdateIndex, + "_BaseUpdateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3809,21 +3780,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_http_options() ) - request, metadata = self._interceptor.pre_upsert_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpsertDatapoints, + "_BaseUpsertDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4026,19 +3994,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4171,15 +4136,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4312,17 +4278,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4456,23 +4421,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4606,17 +4564,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4746,17 +4705,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4861,17 +4821,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4979,17 +4940,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5122,15 +5082,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5263,15 +5224,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/index_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/index_service/transports/rest_asyncio.py index a09c95b7b2..a84ab1aeda 100644 --- a/google/cloud/aiplatform_v1/services/index_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/index_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -971,27 +973,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCreateIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_create_index( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCreateIndex, + "_BaseCreateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1139,21 +1132,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_index( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteIndex, + "_BaseDeleteIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1301,19 +1291,16 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_get_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIndex, + "_BaseGetIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1459,21 +1446,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListIndexes._get_http_options() ) - request, metadata = await self._interceptor.pre_list_indexes( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListIndexes, + "_BaseListIndexes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1622,21 +1606,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_http_options() ) - request, metadata = await self._interceptor.pre_remove_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseRemoveDatapoints, + "_BaseRemoveDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1789,27 +1770,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpdateIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_update_index( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpdateIndex, + "_BaseUpdateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1957,21 +1929,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_http_options() ) - request, metadata = await self._interceptor.pre_upsert_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpsertDatapoints, + "_BaseUpsertDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4161,21 +4130,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4314,17 +4280,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4465,19 +4432,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4617,25 +4583,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4776,17 +4735,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4924,17 +4884,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5047,17 +5008,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5173,19 +5135,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5324,17 +5285,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5475,17 +5437,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/index_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/index_service/transports/rest_base.py index 40336e5c4d..e25308692f 100644 --- a/google/cloud/aiplatform_v1/services/index_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/index_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.CreateIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,43 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.DeleteIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -222,43 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.GetIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseGetIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListIndexes: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -269,43 +159,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.ListIndexesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseListIndexes._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRemoveDatapoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -317,52 +176,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.RemoveDatapointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -374,52 +193,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.UpdateIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpsertDatapoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -431,38 +210,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.UpsertDatapointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -481,17 +228,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -510,17 +246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -591,17 +316,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -682,22 +396,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -764,17 +462,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1149,17 +836,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1562,17 +1238,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1987,17 +1652,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2408,17 +2062,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2825,16 +2468,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseIndexServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/job_service/client.py b/google/cloud/aiplatform_v1/services/job_service/client.py index 21e187a284..646567d865 100644 --- a/google/cloud/aiplatform_v1/services/job_service/client.py +++ b/google/cloud/aiplatform_v1/services/job_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -164,77 +171,13 @@ def get_transport_class( class JobServiceClient(metaclass=JobServiceClientMeta): """A service for creating and managing Vertex AI's jobs.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -799,7 +742,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = JobServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -820,35 +763,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = JobServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -868,65 +788,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = JobServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = JobServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = JobServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -1056,13 +917,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - JobServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = JobServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = JobServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=JobServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -1098,11 +961,14 @@ def __init__( self._transport = cast(JobServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or JobServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=JobServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=JobServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/job_service/transports/rest.py b/google/cloud/aiplatform_v1/services/job_service/transports/rest.py index 3c691c2a8d..46806e8389 100644 --- a/google/cloud/aiplatform_v1/services/job_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/job_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -4246,21 +4247,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob, + "_BaseCancelBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4362,21 +4360,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelCustomJob, + "_BaseCancelCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4478,21 +4473,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelDataLabelingJob, + "_BaseCancelDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4596,21 +4588,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob, + "_BaseCancelHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4714,25 +4703,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelNasJob, + "_BaseCancelNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4846,21 +4826,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob, + "_BaseCreateBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5010,21 +4987,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateCustomJob, + "_BaseCreateCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5167,21 +5141,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateDataLabelingJob, + "_BaseCreateDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5329,21 +5300,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob, + "_BaseCreateHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5498,23 +5466,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_create_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob, + "_BaseCreateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5664,25 +5629,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateNasJob, + "_BaseCreateNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5825,17 +5781,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob, + "_BaseDeleteBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5974,17 +5931,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteCustomJob, + "_BaseDeleteCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6123,17 +6081,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob, + "_BaseDeleteDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6274,17 +6233,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob, + "_BaseDeleteHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6429,19 +6389,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_delete_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob, + "_BaseDeleteModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6582,19 +6543,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteNasJob, + "_BaseDeleteNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6737,17 +6695,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetBatchPredictionJob, + "_BaseGetBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6895,19 +6854,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_custom_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetCustomJob, + "_BaseGetCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7048,17 +7004,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetDataLabelingJob, + "_BaseGetDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7203,17 +7160,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob, + "_BaseGetHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7366,19 +7324,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_get_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob, + "_BaseGetModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7524,19 +7483,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasJob, + "_BaseGetNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7678,17 +7634,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_http_options() ) - request, metadata = self._interceptor.pre_get_nas_trial_detail( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasTrialDetail, + "_BaseGetNasTrialDetail__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7829,17 +7786,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_batch_prediction_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListBatchPredictionJobs, + "_BaseListBatchPredictionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7981,19 +7939,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListCustomJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_custom_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListCustomJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListCustomJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListCustomJobs, + "_BaseListCustomJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8135,17 +8092,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_data_labeling_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListDataLabelingJobs, + "_BaseListDataLabelingJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8289,17 +8247,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_hyperparameter_tuning_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs, + "_BaseListHyperparameterTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8449,19 +8408,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_http_options() ) - request, metadata = ( self._interceptor.pre_list_model_deployment_monitoring_jobs( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs, + "_BaseListModelDeploymentMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8607,19 +8567,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_nas_jobs(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasJobs, + "_BaseListNasJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8759,17 +8716,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_http_options() ) - request, metadata = self._interceptor.pre_list_nas_trial_details( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasTrialDetails, + "_BaseListNasTrialDetails__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8908,23 +8866,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_pause_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob, + "_BasePauseModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9028,23 +8983,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_resume_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob, + "_BaseResumeModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9156,23 +9108,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_http_options() ) - request, metadata = ( self._interceptor.pre_search_model_deployment_monitoring_stats_anomalies( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies, + "_BaseSearchModelDeploymentMonitoringStatsAnomalies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9327,23 +9276,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_update_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob, + "_BaseUpdateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9819,19 +9765,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9964,19 +9907,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10109,19 +10049,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10255,25 +10192,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10407,17 +10335,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10547,17 +10476,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10662,17 +10592,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10780,19 +10711,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10925,17 +10853,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseJobServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11068,19 +10995,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/job_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/job_service/transports/rest_asyncio.py index fc06c8a802..f02554db75 100644 --- a/google/cloud/aiplatform_v1/services/job_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/job_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -2483,21 +2485,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob, + "_BaseCancelBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2605,21 +2604,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelCustomJob, + "_BaseCancelCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2729,21 +2725,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelDataLabelingJob, + "_BaseCancelDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2855,23 +2848,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_cancel_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob, + "_BaseCancelHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2979,27 +2969,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelNasJob, + "_BaseCancelNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3120,21 +3101,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob, + "_BaseCreateBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3292,21 +3270,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateCustomJob, + "_BaseCreateCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3457,21 +3432,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateDataLabelingJob, + "_BaseCreateDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3629,23 +3601,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob, + "_BaseCreateHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3806,23 +3775,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob, + "_BaseCreateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3980,27 +3946,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateNasJob, + "_BaseCreateNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4150,17 +4107,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob, + "_BaseDeleteBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4309,17 +4267,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteCustomJob, + "_BaseDeleteCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4468,17 +4427,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob, + "_BaseDeleteDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4631,19 +4591,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob, + "_BaseDeleteHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4796,19 +4757,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob, + "_BaseDeleteModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4959,21 +4921,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteNasJob, + "_BaseDeleteNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5124,17 +5083,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetBatchPredictionJob, + "_BaseGetBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5292,21 +5252,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_custom_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetCustomJob, + "_BaseGetCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5453,17 +5410,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetDataLabelingJob, + "_BaseGetDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5617,19 +5575,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob, + "_BaseGetHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5786,19 +5745,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob, + "_BaseGetModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5952,21 +5912,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasJob, + "_BaseGetNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6114,17 +6071,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_http_options() ) - request, metadata = await self._interceptor.pre_get_nas_trial_detail( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasTrialDetail, + "_BaseGetNasTrialDetail__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6274,17 +6232,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_batch_prediction_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListBatchPredictionJobs, + "_BaseListBatchPredictionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6434,19 +6393,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListCustomJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_custom_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListCustomJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListCustomJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListCustomJobs, + "_BaseListCustomJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6594,17 +6552,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_data_labeling_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListDataLabelingJobs, + "_BaseListDataLabelingJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6758,19 +6717,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_hyperparameter_tuning_jobs( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs, + "_BaseListHyperparameterTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6926,19 +6886,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_model_deployment_monitoring_jobs( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs, + "_BaseListModelDeploymentMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7092,21 +7053,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_nas_jobs( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasJobs, + "_BaseListNasJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7252,17 +7210,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_http_options() ) - request, metadata = await self._interceptor.pre_list_nas_trial_details( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasTrialDetails, + "_BaseListNasTrialDetails__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7411,23 +7370,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_pause_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob, + "_BasePauseModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7539,23 +7495,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_resume_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob, + "_BaseResumeModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7673,23 +7626,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_http_options() ) - request, metadata = ( await self._interceptor.pre_search_model_deployment_monitoring_stats_anomalies( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies, + "_BaseSearchModelDeploymentMonitoringStatsAnomalies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7852,23 +7802,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob, + "_BaseUpdateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10277,21 +10224,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10430,21 +10374,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10583,21 +10524,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10737,27 +10675,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10897,17 +10826,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11045,17 +10975,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11168,17 +11099,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11294,21 +11226,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11447,19 +11376,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11598,21 +11526,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/job_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/job_service/transports/rest_base.py index 58bcca2184..bc06b8d9af 100644 --- a/google/cloud/aiplatform_v1/services/job_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/job_service/transports/rest_base.py @@ -118,14 +118,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -137,52 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -194,52 +146,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -251,52 +163,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -308,52 +180,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -365,52 +197,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateBatchPredictionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -422,52 +214,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -479,52 +231,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -536,52 +248,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -593,52 +265,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -650,54 +282,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -709,52 +299,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteBatchPredictionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -765,43 +315,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -812,43 +331,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -859,43 +347,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -906,43 +363,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -953,45 +379,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1002,43 +395,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetBatchPredictionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1049,43 +411,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1096,43 +427,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1143,43 +443,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1190,43 +459,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1237,43 +475,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetModelDeploymentMonitoringJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1284,43 +491,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNasTrialDetail: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1331,43 +507,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetNasTrialDetailRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListBatchPredictionJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1378,43 +523,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListBatchPredictionJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListCustomJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1425,43 +539,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListCustomJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListCustomJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDataLabelingJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1472,43 +555,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListDataLabelingJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListHyperparameterTuningJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1519,43 +571,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListHyperparameterTuningJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelDeploymentMonitoringJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1566,45 +587,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListModelDeploymentMonitoringJobsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNasJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1615,43 +603,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListNasJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListNasJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNasTrialDetails: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1662,43 +619,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListNasTrialDetailsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePauseModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1710,54 +636,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.PauseModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseResumeModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1769,54 +653,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ResumeModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchModelDeploymentMonitoringStatsAnomalies: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1828,42 +670,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - job_service.SearchModelDeploymentMonitoringStatsAnomaliesRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1872,14 +678,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1891,40 +689,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.UpdateModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1943,17 +707,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1972,17 +725,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2053,17 +795,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2144,22 +875,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2226,17 +941,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2611,17 +1315,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3024,17 +1717,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3449,17 +2131,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3870,17 +2541,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4287,16 +2947,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseJobServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/llm_utility_service/client.py b/google/cloud/aiplatform_v1/services/llm_utility_service/client.py index 56fa57653d..6fe3d1ee49 100644 --- a/google/cloud/aiplatform_v1/services/llm_utility_service/client.py +++ b/google/cloud/aiplatform_v1/services/llm_utility_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -130,77 +137,13 @@ def get_transport_class( class LlmUtilityServiceClient(metaclass=LlmUtilityServiceClientMeta): """Service for LLM related utility functions.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -410,7 +353,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = LlmUtilityServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -431,35 +374,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = LlmUtilityServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -479,65 +399,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = LlmUtilityServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = LlmUtilityServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -671,13 +532,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - LlmUtilityServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = LlmUtilityServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = LlmUtilityServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=LlmUtilityServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -713,14 +576,14 @@ def __init__( self._transport = cast(LlmUtilityServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or LlmUtilityServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=LlmUtilityServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest.py b/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest.py index f0015b06b9..949c2bca6b 100644 --- a/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -606,19 +607,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_http_options() ) - request, metadata = self._interceptor.pre_compute_tokens(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseComputeTokens, + "_BaseComputeTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -759,19 +757,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_http_options() ) - request, metadata = self._interceptor.pre_count_tokens(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseCountTokens, + "_BaseCountTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -935,15 +930,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1077,15 +1073,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1218,15 +1215,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1360,19 +1358,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1507,17 +1502,18 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1648,17 +1644,18 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1764,17 +1761,18 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1882,15 +1880,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2024,15 +2023,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2166,15 +2166,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_asyncio.py index 1de4a0997e..c2039367c7 100644 --- a/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -666,21 +668,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_http_options() ) - request, metadata = await self._interceptor.pre_compute_tokens( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseComputeTokens, + "_BaseComputeTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -830,21 +829,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_http_options() ) - request, metadata = await self._interceptor.pre_count_tokens( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseCountTokens, + "_BaseCountTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1013,17 +1009,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1165,17 +1162,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1317,17 +1315,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1470,21 +1469,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1627,17 +1623,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1774,17 +1771,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1896,17 +1894,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2021,17 +2020,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2173,17 +2173,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2325,17 +2326,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_base.py index e77942f5e8..a07bcf20b4 100644 --- a/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/llm_utility_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -134,52 +126,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = llm_utility_service.ComputeTokensRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCountTokens: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -206,38 +158,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.CountTokensRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseLlmUtilityServiceRestTransport._BaseCountTokens._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -256,17 +176,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -285,17 +194,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -366,17 +264,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -457,22 +344,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -539,17 +410,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -924,17 +784,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1337,17 +1186,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1762,17 +1600,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2183,17 +2010,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2600,16 +2416,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseLlmUtilityServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/match_service/client.py b/google/cloud/aiplatform_v1/services/match_service/client.py index 05aff95f8f..1d676bbf0b 100644 --- a/google/cloud/aiplatform_v1/services/match_service/client.py +++ b/google/cloud/aiplatform_v1/services/match_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -128,77 +135,13 @@ class MatchServiceClient(metaclass=MatchServiceClientMeta): similarity search at scale. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -386,7 +329,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = MatchServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -407,35 +350,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = MatchServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -455,65 +375,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = MatchServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = MatchServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = MatchServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -643,13 +504,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - MatchServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = MatchServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = MatchServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=MatchServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -685,11 +548,14 @@ def __init__( self._transport = cast(MatchServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or MatchServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=MatchServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=MatchServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/match_service/transports/rest.py b/google/cloud/aiplatform_v1/services/match_service/transports/rest.py index 85757f577e..020af8eea5 100644 --- a/google/cloud/aiplatform_v1/services/match_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/match_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -605,19 +606,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseFindNeighbors._get_http_options() ) - request, metadata = self._interceptor.pre_find_neighbors(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseFindNeighbors, + "_BaseFindNeighbors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -761,21 +759,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_http_options() ) - request, metadata = self._interceptor.pre_read_index_datapoints( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseReadIndexDatapoints, + "_BaseReadIndexDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -939,19 +934,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1084,15 +1076,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1225,17 +1218,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1369,23 +1361,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1519,17 +1504,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1659,17 +1645,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1774,17 +1761,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1892,17 +1880,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2035,15 +2022,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2176,15 +2164,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/match_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/match_service/transports/rest_asyncio.py index eeb763641b..d1ceb39866 100644 --- a/google/cloud/aiplatform_v1/services/match_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/match_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -665,21 +667,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseFindNeighbors._get_http_options() ) - request, metadata = await self._interceptor.pre_find_neighbors( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseFindNeighbors, + "_BaseFindNeighbors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -832,21 +831,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_http_options() ) - request, metadata = await self._interceptor.pre_read_index_datapoints( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseReadIndexDatapoints, + "_BaseReadIndexDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1014,21 +1010,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1167,17 +1160,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1318,19 +1312,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1470,25 +1463,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1629,17 +1615,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1777,17 +1764,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1900,17 +1888,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2026,19 +2015,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2177,17 +2165,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2328,17 +2317,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/match_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/match_service/transports/rest_base.py index 14cd46274b..39a75e3029 100644 --- a/google/cloud/aiplatform_v1/services/match_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/match_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,52 +110,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = match_service.FindNeighborsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMatchServiceRestTransport._BaseFindNeighbors._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadIndexDatapoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,38 +127,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = match_service.ReadIndexDatapointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -225,17 +145,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -254,17 +163,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -335,17 +233,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -426,22 +313,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -508,17 +379,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -893,17 +753,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1306,17 +1155,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1731,17 +1569,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2152,17 +1979,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2569,16 +2385,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseMatchServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/metadata_service/client.py b/google/cloud/aiplatform_v1/services/metadata_service/client.py index 09b1393ba7..1b841f1c2f 100644 --- a/google/cloud/aiplatform_v1/services/metadata_service/client.py +++ b/google/cloud/aiplatform_v1/services/metadata_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -148,77 +155,13 @@ def get_transport_class( class MetadataServiceClient(metaclass=MetadataServiceClientMeta): """Service for reading and writing metadata entries.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -502,7 +445,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = MetadataServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -523,35 +466,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = MetadataServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -571,65 +491,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = MetadataServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = MetadataServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -761,13 +622,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - MetadataServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = MetadataServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = MetadataServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=MetadataServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -803,14 +666,14 @@ def __init__( self._transport = cast(MetadataServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or MetadataServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=MetadataServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=MetadataServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/metadata_service/transports/rest.py b/google/cloud/aiplatform_v1/services/metadata_service/transports/rest.py index b7544961ab..1651278faf 100644 --- a/google/cloud/aiplatform_v1/services/metadata_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/metadata_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -4322,23 +4323,20 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_http_options() ) - request, metadata = ( self._interceptor.pre_add_context_artifacts_and_executions( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions, + "_BaseAddContextArtifactsAndExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4485,21 +4483,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_http_options() ) - request, metadata = self._interceptor.pre_add_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextChildren, + "_BaseAddContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4644,21 +4639,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_http_options() ) - request, metadata = self._interceptor.pre_add_execution_events( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddExecutionEvents, + "_BaseAddExecutionEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4800,19 +4792,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_create_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateArtifact, + "_BaseCreateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4952,19 +4941,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateContext._get_http_options() ) - request, metadata = self._interceptor.pre_create_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateContext, + "_BaseCreateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5104,21 +5090,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateExecution._get_http_options() ) - request, metadata = self._interceptor.pre_create_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateExecution, + "_BaseCreateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5259,21 +5242,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_http_options() ) - request, metadata = self._interceptor.pre_create_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema, + "_BaseCreateMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5419,21 +5399,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_http_options() ) - request, metadata = self._interceptor.pre_create_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataStore, + "_BaseCreateMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5573,15 +5550,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_delete_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteArtifact, + "_BaseDeleteArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5720,15 +5698,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteContext._get_http_options() ) - request, metadata = self._interceptor.pre_delete_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteContext, + "_BaseDeleteContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5867,17 +5846,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_http_options() ) - request, metadata = self._interceptor.pre_delete_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteExecution, + "_BaseDeleteExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6017,17 +5997,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_http_options() ) - request, metadata = self._interceptor.pre_delete_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore, + "_BaseDeleteMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6163,15 +6144,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_get_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetArtifact, + "_BaseGetArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6309,15 +6291,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetContext._get_http_options() ) - request, metadata = self._interceptor.pre_get_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetContext, + "_BaseGetContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6455,15 +6438,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetExecution._get_http_options() ) - request, metadata = self._interceptor.pre_get_execution(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetExecution, + "_BaseGetExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6602,17 +6586,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_http_options() ) - request, metadata = self._interceptor.pre_get_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataSchema, + "_BaseGetMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6753,17 +6738,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_http_options() ) - request, metadata = self._interceptor.pre_get_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataStore, + "_BaseGetMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6903,15 +6889,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListArtifacts._get_http_options() ) - request, metadata = self._interceptor.pre_list_artifacts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListArtifacts, + "_BaseListArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7053,15 +7040,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListContexts._get_http_options() ) - request, metadata = self._interceptor.pre_list_contexts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListContexts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListContexts, + "_BaseListContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7203,15 +7191,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListExecutions._get_http_options() ) - request, metadata = self._interceptor.pre_list_executions(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListExecutions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListExecutions, + "_BaseListExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7354,17 +7343,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_http_options() ) - request, metadata = self._interceptor.pre_list_metadata_schemas( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataSchemas, + "_BaseListMetadataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7507,17 +7497,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_http_options() ) - request, metadata = self._interceptor.pre_list_metadata_stores( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataStores, + "_BaseListMetadataStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7661,19 +7652,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_http_options() ) - request, metadata = self._interceptor.pre_purge_artifacts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeArtifacts, + "_BasePurgeArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7814,19 +7802,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeContexts._get_http_options() ) - request, metadata = self._interceptor.pre_purge_contexts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeContexts, + "_BasePurgeContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7967,21 +7952,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_http_options() ) - request, metadata = self._interceptor.pre_purge_executions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeExecutions, + "_BasePurgeExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8123,17 +8105,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_http_options() ) - request, metadata = self._interceptor.pre_query_artifact_lineage_subgraph( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph, + "_BaseQueryArtifactLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8280,17 +8263,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_http_options() ) - request, metadata = self._interceptor.pre_query_context_lineage_subgraph( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph, + "_BaseQueryContextLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8439,19 +8423,20 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_http_options() ) - request, metadata = ( self._interceptor.pre_query_execution_inputs_and_outputs( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs, + "_BaseQueryExecutionInputsAndOutputs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8597,21 +8582,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_http_options() ) - request, metadata = self._interceptor.pre_remove_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseRemoveContextChildren, + "_BaseRemoveContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8755,19 +8737,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_update_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateArtifact, + "_BaseUpdateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8907,19 +8886,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateContext._get_http_options() ) - request, metadata = self._interceptor.pre_update_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateContext, + "_BaseUpdateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9059,21 +9035,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_http_options() ) - request, metadata = self._interceptor.pre_update_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateExecution, + "_BaseUpdateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9515,15 +9488,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9656,15 +9630,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9797,15 +9772,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9939,19 +9915,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10086,17 +10059,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10226,17 +10200,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10341,17 +10316,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10459,15 +10435,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10600,15 +10577,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10741,15 +10719,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_asyncio.py index 28eeabc22d..4426d9e15a 100644 --- a/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -2549,23 +2551,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_http_options() ) - request, metadata = ( await self._interceptor.pre_add_context_artifacts_and_executions( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions, + "_BaseAddContextArtifactsAndExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2720,21 +2719,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_http_options() ) - request, metadata = await self._interceptor.pre_add_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextChildren, + "_BaseAddContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2885,21 +2881,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_http_options() ) - request, metadata = await self._interceptor.pre_add_execution_events( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddExecutionEvents, + "_BaseAddExecutionEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3048,21 +3041,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_create_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateArtifact, + "_BaseCreateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3211,21 +3201,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateContext._get_http_options() ) - request, metadata = await self._interceptor.pre_create_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateContext, + "_BaseCreateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3374,21 +3361,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_create_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateExecution, + "_BaseCreateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3537,21 +3521,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_http_options() ) - request, metadata = await self._interceptor.pre_create_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema, + "_BaseCreateMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3703,21 +3684,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_http_options() ) - request, metadata = await self._interceptor.pre_create_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataStore, + "_BaseCreateMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3866,17 +3844,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteArtifact, + "_BaseDeleteArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4026,17 +4005,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteContext._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteContext, + "_BaseDeleteContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4186,17 +4166,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteExecution, + "_BaseDeleteExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4346,17 +4327,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore, + "_BaseDeleteMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4500,17 +4482,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_get_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetArtifact, + "_BaseGetArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4656,17 +4639,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetContext._get_http_options() ) - request, metadata = await self._interceptor.pre_get_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetContext, + "_BaseGetContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4813,17 +4797,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_get_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetExecution, + "_BaseGetExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4970,17 +4955,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_http_options() ) - request, metadata = await self._interceptor.pre_get_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataSchema, + "_BaseGetMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5128,17 +5114,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_http_options() ) - request, metadata = await self._interceptor.pre_get_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataStore, + "_BaseGetMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5287,17 +5274,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListArtifacts._get_http_options() ) - request, metadata = await self._interceptor.pre_list_artifacts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListArtifacts, + "_BaseListArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5448,17 +5436,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_list_contexts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListContexts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListContexts, + "_BaseListContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5609,17 +5598,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListExecutions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_executions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListExecutions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListExecutions, + "_BaseListExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5770,17 +5760,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_http_options() ) - request, metadata = await self._interceptor.pre_list_metadata_schemas( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataSchemas, + "_BaseListMetadataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5929,17 +5920,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_http_options() ) - request, metadata = await self._interceptor.pre_list_metadata_stores( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataStores, + "_BaseListMetadataStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6090,21 +6082,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_http_options() ) - request, metadata = await self._interceptor.pre_purge_artifacts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeArtifacts, + "_BasePurgeArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6256,21 +6245,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_purge_contexts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeContexts, + "_BasePurgeContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6422,21 +6408,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_http_options() ) - request, metadata = await self._interceptor.pre_purge_executions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeExecutions, + "_BasePurgeExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6590,19 +6573,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_http_options() ) - request, metadata = ( await self._interceptor.pre_query_artifact_lineage_subgraph( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph, + "_BaseQueryArtifactLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6755,19 +6739,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_http_options() ) - request, metadata = ( await self._interceptor.pre_query_context_lineage_subgraph( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph, + "_BaseQueryContextLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6922,19 +6907,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_query_execution_inputs_and_outputs( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs, + "_BaseQueryExecutionInputsAndOutputs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7086,21 +7072,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_http_options() ) - request, metadata = await self._interceptor.pre_remove_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseRemoveContextChildren, + "_BaseRemoveContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7251,21 +7234,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_update_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateArtifact, + "_BaseUpdateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7414,21 +7394,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateContext._get_http_options() ) - request, metadata = await self._interceptor.pre_update_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateContext, + "_BaseUpdateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7577,21 +7554,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_update_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateExecution, + "_BaseUpdateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9970,17 +9944,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10122,17 +10097,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10274,17 +10250,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10427,21 +10404,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10584,17 +10558,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10731,17 +10706,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10855,17 +10831,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10982,17 +10959,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11134,17 +11112,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11286,17 +11265,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_base.py index 84812158e5..b12c39d2fd 100644 --- a/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/metadata_service/transports/rest_base.py @@ -109,14 +109,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -128,54 +120,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.AddContextArtifactsAndExecutionsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAddContextChildren: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -187,52 +137,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.AddContextChildrenRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAddExecutionEvents: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -244,52 +154,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.AddExecutionEventsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -301,52 +171,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -358,52 +188,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -415,52 +205,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateMetadataSchema: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -472,52 +222,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateMetadataSchemaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateMetadataStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -529,52 +239,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateMetadataStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -585,43 +255,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -632,43 +271,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -679,43 +287,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteMetadataStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -726,43 +303,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteMetadataStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -773,43 +319,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -820,43 +335,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -867,43 +351,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetMetadataSchema: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -914,43 +367,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetMetadataSchemaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetMetadataStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -961,43 +383,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetMetadataStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListArtifacts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1008,43 +399,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListArtifactsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListArtifacts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1055,43 +415,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListExecutions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1102,43 +431,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListExecutionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListExecutions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListMetadataSchemas: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1149,43 +447,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListMetadataSchemasRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListMetadataStores: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1196,43 +463,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListMetadataStoresRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePurgeArtifacts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1244,52 +480,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.PurgeArtifactsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePurgeContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1301,52 +497,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.PurgeContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BasePurgeContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePurgeExecutions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1358,52 +514,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.PurgeExecutionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryArtifactLineageSubgraph: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1414,45 +530,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.QueryArtifactLineageSubgraphRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryContextLineageSubgraph: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1463,43 +546,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.QueryContextLineageSubgraphRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryExecutionInputsAndOutputs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1510,45 +562,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.QueryExecutionInputsAndOutputsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRemoveContextChildren: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1560,52 +579,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.RemoveContextChildrenRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1617,52 +596,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.UpdateArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1674,52 +613,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.UpdateContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseUpdateContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1731,38 +630,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.UpdateExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1781,17 +648,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1810,17 +666,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1891,17 +736,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1982,22 +816,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2064,17 +882,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2449,17 +1256,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2862,17 +1658,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3287,17 +2072,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3708,17 +2482,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4125,16 +2888,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseMetadataServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/migration_service/client.py b/google/cloud/aiplatform_v1/services/migration_service/client.py index 4d5443989f..cb18ce0ea6 100644 --- a/google/cloud/aiplatform_v1/services/migration_service/client.py +++ b/google/cloud/aiplatform_v1/services/migration_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -133,77 +140,13 @@ class MigrationServiceClient(metaclass=MigrationServiceClientMeta): datalabeling.googleapis.com and ml.googleapis.com to Vertex AI. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -518,7 +461,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = MigrationServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -539,35 +482,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = MigrationServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -587,65 +507,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = MigrationServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = MigrationServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -777,13 +638,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - MigrationServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = MigrationServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = MigrationServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=MigrationServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -819,14 +682,14 @@ def __init__( self._transport = cast(MigrationServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or MigrationServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=MigrationServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=MigrationServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/migration_service/transports/rest.py b/google/cloud/aiplatform_v1/services/migration_service/transports/rest.py index 795a0f96f2..1bd6526a08 100644 --- a/google/cloud/aiplatform_v1/services/migration_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/migration_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2602,21 +2603,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_http_options() ) - request, metadata = self._interceptor.pre_batch_migrate_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseBatchMigrateResources, + "_BaseBatchMigrateResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2760,21 +2758,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_http_options() ) - request, metadata = self._interceptor.pre_search_migratable_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSearchMigratableResources, + "_BaseSearchMigratableResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2942,15 +2937,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3083,15 +3079,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3224,15 +3221,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3366,19 +3364,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3513,17 +3508,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3654,17 +3650,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3770,17 +3767,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3888,15 +3886,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4029,15 +4028,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4170,15 +4170,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/migration_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/migration_service/transports/rest_asyncio.py index 49e01dd417..918ad49ace 100644 --- a/google/cloud/aiplatform_v1/services/migration_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/migration_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -671,21 +673,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_migrate_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseBatchMigrateResources, + "_BaseBatchMigrateResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -837,21 +836,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_http_options() ) - request, metadata = await self._interceptor.pre_search_migratable_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSearchMigratableResources, + "_BaseSearchMigratableResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3015,17 +3011,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3167,17 +3164,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3319,17 +3317,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3472,21 +3471,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3629,17 +3625,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3776,17 +3773,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3900,17 +3898,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4027,17 +4026,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4179,17 +4179,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4331,17 +4332,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/migration_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/migration_service/transports/rest_base.py index 6d3ff7ffee..cb1e80e002 100644 --- a/google/cloud/aiplatform_v1/services/migration_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/migration_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,52 +110,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = migration_service.BatchMigrateResourcesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchMigratableResources: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,38 +127,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = migration_service.SearchMigratableResourcesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -225,17 +145,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -254,17 +163,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -335,17 +233,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -426,22 +313,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -508,17 +379,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -893,17 +753,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1306,17 +1155,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1731,17 +1569,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2152,17 +1979,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2569,16 +2385,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseMigrationServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/model_garden_service/client.py b/google/cloud/aiplatform_v1/services/model_garden_service/client.py index 6a43b34703..00389df2ba 100644 --- a/google/cloud/aiplatform_v1/services/model_garden_service/client.py +++ b/google/cloud/aiplatform_v1/services/model_garden_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -131,77 +138,13 @@ def get_transport_class( class ModelGardenServiceClient(metaclass=ModelGardenServiceClientMeta): """The interface of Model Garden Service.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -450,7 +393,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ModelGardenServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -471,35 +414,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ModelGardenServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -519,65 +439,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ModelGardenServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ModelGardenServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -711,13 +572,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ModelGardenServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ModelGardenServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ModelGardenServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ModelGardenServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -753,14 +616,14 @@ def __init__( self._transport = cast(ModelGardenServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ModelGardenServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ModelGardenServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest.py b/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest.py index 1108842051..418c9098fc 100644 --- a/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2597,23 +2598,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeploy._get_http_options() ) - request, metadata = self._interceptor.pre_deploy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeploy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeploy, + "_BaseDeploy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2751,17 +2745,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_http_options() ) - request, metadata = self._interceptor.pre_get_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetPublisherModel, + "_BaseGetPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2920,15 +2915,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3062,15 +3058,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3204,15 +3201,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3347,19 +3345,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3494,17 +3489,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3637,17 +3633,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3753,17 +3750,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3872,15 +3870,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4014,15 +4013,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4156,15 +4156,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_asyncio.py index f6f08bc4ee..de9e82b9e7 100644 --- a/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -667,23 +669,16 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeploy._get_http_options() ) - request, metadata = await self._interceptor.pre_deploy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeploy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeploy, + "_BaseDeploy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -829,17 +824,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_http_options() ) - request, metadata = await self._interceptor.pre_get_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetPublisherModel, + "_BaseGetPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2993,17 +2989,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3145,17 +3142,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3297,17 +3295,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3450,21 +3449,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3607,17 +3603,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3754,17 +3751,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3876,17 +3874,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4001,17 +4000,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4153,17 +4153,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4303,17 +4304,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_base.py index 28b153a34d..ecbdf5e03c 100644 --- a/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/model_garden_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.DeployRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetPublisherModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,29 +127,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.GetPublisherModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -216,17 +145,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -245,17 +163,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -326,17 +233,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -417,22 +313,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -499,17 +379,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -884,17 +753,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1297,17 +1155,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1722,17 +1569,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2143,17 +1979,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2560,16 +2385,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseModelGardenServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/model_service/client.py b/google/cloud/aiplatform_v1/services/model_service/client.py index 5abede94bf..3c2230509d 100644 --- a/google/cloud/aiplatform_v1/services/model_service/client.py +++ b/google/cloud/aiplatform_v1/services/model_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -142,77 +149,13 @@ def get_transport_class( class ModelServiceClient(metaclass=ModelServiceClientMeta): """A service for managing Vertex AI's machine learning Models.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -516,7 +459,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ModelServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -537,35 +480,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ModelServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -585,65 +505,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ModelServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ModelServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ModelServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -773,13 +634,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ModelServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ModelServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ModelServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ModelServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -815,11 +678,14 @@ def __init__( self._transport = cast(ModelServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or ModelServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ModelServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ModelServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/model_service/transports/rest.py b/google/cloud/aiplatform_v1/services/model_service/transports/rest.py index 4d6f5d7f77..256c7d2b2c 100644 --- a/google/cloud/aiplatform_v1/services/model_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/model_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3574,23 +3575,20 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_import_evaluated_annotations( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations, + "_BaseBatchImportEvaluatedAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3740,23 +3738,20 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_import_model_evaluation_slices( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices, + "_BaseBatchImportModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3905,23 +3900,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCopyModel._get_http_options() ) - request, metadata = self._interceptor.pre_copy_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_transcoded_request( - http_options, request - ) - ) - - body = _BaseModelServiceRestTransport._BaseCopyModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCopyModel, + "_BaseCopyModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4061,19 +4049,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModel._get_http_options() ) - request, metadata = self._interceptor.pre_delete_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModel, + "_BaseDeleteModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4212,17 +4197,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_http_options() ) - request, metadata = self._interceptor.pre_delete_model_version( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModelVersion, + "_BaseDeleteModelVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4362,25 +4348,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseExportModel._get_http_options() ) - request, metadata = self._interceptor.pre_export_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseExportModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseExportModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseExportModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseExportModel, + "_BaseExportModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4515,19 +4492,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModel._get_http_options() ) - request, metadata = self._interceptor.pre_get_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModel, + "_BaseGetModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4669,17 +4643,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_http_options() ) - request, metadata = self._interceptor.pre_get_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluation, + "_BaseGetModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4825,17 +4800,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_http_options() ) - request, metadata = self._interceptor.pre_get_model_evaluation_slice( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice, + "_BaseGetModelEvaluationSlice__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4980,21 +4956,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_http_options() ) - request, metadata = self._interceptor.pre_import_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseImportModelEvaluation, + "_BaseImportModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5137,17 +5110,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluations._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_evaluations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluations, + "_BaseListModelEvaluations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5291,17 +5265,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_evaluation_slices( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluationSlices, + "_BaseListModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5447,19 +5422,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModels._get_http_options() ) - request, metadata = self._interceptor.pre_list_models(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseListModels._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseListModels._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModels, + "_BaseListModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5603,17 +5575,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_version_checkpoints( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints, + "_BaseListModelVersionCheckpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5761,17 +5734,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersions._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_versions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersions, + "_BaseListModelVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5912,21 +5886,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_http_options() ) - request, metadata = self._interceptor.pre_merge_version_aliases( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseMergeVersionAliases, + "_BaseMergeVersionAliases__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6071,21 +6042,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_http_options() ) - request, metadata = self._interceptor.pre_update_explanation_dataset( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateExplanationDataset, + "_BaseUpdateExplanationDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6225,25 +6193,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateModel._get_http_options() ) - request, metadata = self._interceptor.pre_update_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateModel, + "_BaseUpdateModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6386,25 +6345,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUploadModel._get_http_options() ) - request, metadata = self._interceptor.pre_upload_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUploadModel, + "_BaseUploadModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6721,19 +6671,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6866,15 +6813,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7007,17 +6955,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7151,23 +7098,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7301,17 +7241,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7441,17 +7382,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7556,17 +7498,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7674,17 +7617,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7817,15 +7759,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7958,15 +7901,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/model_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/model_service/transports/rest_asyncio.py index f00e0f330c..ff3424db18 100644 --- a/google/cloud/aiplatform_v1/services/model_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/model_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1730,23 +1732,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_import_evaluated_annotations( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations, + "_BaseBatchImportEvaluatedAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1904,23 +1903,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_import_model_evaluation_slices( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices, + "_BaseBatchImportModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2077,25 +2073,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCopyModel._get_http_options() ) - request, metadata = await self._interceptor.pre_copy_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_transcoded_request( - http_options, request - ) - ) - - body = _BaseModelServiceRestTransport._BaseCopyModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCopyModel, + "_BaseCopyModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2243,21 +2232,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModel._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModel, + "_BaseDeleteModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2405,17 +2391,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_model_version( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModelVersion, + "_BaseDeleteModelVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2565,27 +2552,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseExportModel._get_http_options() ) - request, metadata = await self._interceptor.pre_export_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseExportModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseExportModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseExportModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseExportModel, + "_BaseExportModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2730,19 +2708,16 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModel._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModel, + "_BaseGetModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2891,17 +2866,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluation, + "_BaseGetModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3055,17 +3031,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model_evaluation_slice( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice, + "_BaseGetModelEvaluationSlice__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3219,21 +3196,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_http_options() ) - request, metadata = await self._interceptor.pre_import_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseImportModelEvaluation, + "_BaseImportModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3385,17 +3359,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_model_evaluations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluations, + "_BaseListModelEvaluations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3545,19 +3520,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_model_evaluation_slices( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluationSlices, + "_BaseListModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3709,21 +3685,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModels._get_http_options() ) - request, metadata = await self._interceptor.pre_list_models( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseListModels._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseListModels._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModels, + "_BaseListModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3873,19 +3846,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_model_version_checkpoints( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints, + "_BaseListModelVersionCheckpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4037,17 +4011,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_model_versions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersions, + "_BaseListModelVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4197,21 +4172,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_http_options() ) - request, metadata = await self._interceptor.pre_merge_version_aliases( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseMergeVersionAliases, + "_BaseMergeVersionAliases__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4364,21 +4336,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_update_explanation_dataset( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateExplanationDataset, + "_BaseUpdateExplanationDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4526,27 +4495,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateModel._get_http_options() ) - request, metadata = await self._interceptor.pre_update_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateModel, + "_BaseUpdateModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4695,27 +4655,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUploadModel._get_http_options() ) - request, metadata = await self._interceptor.pre_upload_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUploadModel, + "_BaseUploadModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6995,21 +6946,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7148,17 +7096,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7299,19 +7248,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7451,25 +7399,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7610,17 +7551,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7758,17 +7700,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7881,17 +7824,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8007,19 +7951,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8158,17 +8101,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8309,17 +8253,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/model_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/model_service/transports/rest_base.py index da11cf742e..95bc897961 100644 --- a/google/cloud/aiplatform_v1/services/model_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/model_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,54 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.BatchImportEvaluatedAnnotationsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchImportModelEvaluationSlices: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -182,54 +132,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.BatchImportModelEvaluationSlicesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCopyModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -241,52 +149,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.CopyModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseCopyModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -297,43 +165,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.DeleteModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseDeleteModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModelVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -344,43 +181,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.DeleteModelVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -392,52 +198,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ExportModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseExportModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -448,43 +214,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.GetModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseGetModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelEvaluation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -495,43 +230,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.GetModelEvaluationRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelEvaluationSlice: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -542,43 +246,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.GetModelEvaluationSliceRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportModelEvaluation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -590,52 +263,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ImportModelEvaluationRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelEvaluations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -646,43 +279,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelEvaluationsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelEvaluations._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelEvaluationSlices: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -693,43 +295,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelEvaluationSlicesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModels: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -740,43 +311,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModels._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelVersionCheckpoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -787,43 +327,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelVersionCheckpointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelVersions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -834,43 +343,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelVersionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelVersions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseMergeVersionAliases: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -882,52 +360,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.MergeVersionAliasesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateExplanationDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -939,38 +377,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.UpdateExplanationDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -979,14 +385,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -998,52 +396,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.UpdateModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseUpdateModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUploadModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1055,38 +413,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.UploadModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseUploadModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1105,17 +431,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1134,17 +449,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1215,17 +519,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1306,22 +599,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1388,17 +665,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1773,17 +1039,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2186,17 +1441,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2611,17 +1855,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3032,17 +2265,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3449,16 +2671,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseModelServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/notebook_service/client.py b/google/cloud/aiplatform_v1/services/notebook_service/client.py index 3e913d231c..929dc071d1 100644 --- a/google/cloud/aiplatform_v1/services/notebook_service/client.py +++ b/google/cloud/aiplatform_v1/services/notebook_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -152,77 +159,13 @@ class NotebookServiceClient(metaclass=NotebookServiceClientMeta): Workbench). """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -539,7 +482,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = NotebookServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -560,35 +503,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = NotebookServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -608,65 +528,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = NotebookServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = NotebookServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -798,13 +659,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - NotebookServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = NotebookServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = NotebookServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=NotebookServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -840,14 +703,14 @@ def __init__( self._transport = cast(NotebookServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or NotebookServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=NotebookServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=NotebookServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/notebook_service/transports/rest.py b/google/cloud/aiplatform_v1/services/notebook_service/transports/rest.py index 14c8046922..59fc2dbfd9 100644 --- a/google/cloud/aiplatform_v1/services/notebook_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/notebook_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3419,21 +3420,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_assign_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime, + "_BaseAssignNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3578,21 +3576,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob, + "_BaseCreateNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3739,21 +3734,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_create_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate, + "_BaseCreateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3897,17 +3889,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob, + "_BaseDeleteNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4051,17 +4044,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_delete_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime, + "_BaseDeleteNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4204,17 +4198,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_delete_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate, + "_BaseDeleteNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4356,17 +4351,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob, + "_BaseGetNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4516,17 +4512,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_get_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime, + "_BaseGetNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4673,17 +4670,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_get_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate, + "_BaseGetNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4831,17 +4829,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_notebook_execution_jobs( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs, + "_BaseListNotebookExecutionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4988,17 +4987,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_http_options() ) - request, metadata = self._interceptor.pre_list_notebook_runtimes( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes, + "_BaseListNotebookRuntimes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5142,17 +5142,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_http_options() ) - request, metadata = self._interceptor.pre_list_notebook_runtime_templates( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates, + "_BaseListNotebookRuntimeTemplates__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5301,21 +5302,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_start_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime, + "_BaseStartNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5457,21 +5455,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_stop_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime, + "_BaseStopNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5616,21 +5611,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_update_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate, + "_BaseUpdateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5778,21 +5770,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_upgrade_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime, + "_BaseUpgradeNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6101,15 +6090,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6242,15 +6232,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6383,15 +6374,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6525,19 +6517,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6672,17 +6661,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6812,17 +6802,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6927,17 +6918,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7045,15 +7037,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7186,15 +7179,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7327,15 +7321,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_asyncio.py index a4e12b9e20..b04fba4562 100644 --- a/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1558,21 +1560,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_assign_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime, + "_BaseAssignNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1725,23 +1724,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_notebook_execution_job( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob, + "_BaseCreateNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1896,23 +1892,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate, + "_BaseCreateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2064,19 +2057,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_notebook_execution_job( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob, + "_BaseDeleteNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2226,17 +2220,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime, + "_BaseDeleteNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2389,19 +2384,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate, + "_BaseDeleteNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2551,17 +2547,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob, + "_BaseGetNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2717,17 +2714,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_get_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime, + "_BaseGetNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2880,19 +2878,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate, + "_BaseGetNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3044,19 +3043,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_notebook_execution_jobs( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs, + "_BaseListNotebookExecutionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3209,17 +3209,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_http_options() ) - request, metadata = await self._interceptor.pre_list_notebook_runtimes( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes, + "_BaseListNotebookRuntimes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3371,19 +3372,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_notebook_runtime_templates( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates, + "_BaseListNotebookRuntimeTemplates__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3538,21 +3540,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_start_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime, + "_BaseStartNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3702,21 +3701,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_stop_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime, + "_BaseStopNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3871,23 +3867,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate, + "_BaseUpdateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4041,21 +4034,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_upgrade_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime, + "_BaseUpgradeNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6333,17 +6323,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6485,17 +6476,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6637,17 +6629,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6790,21 +6783,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6947,17 +6937,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7094,17 +7085,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7218,17 +7210,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7345,17 +7338,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7497,17 +7491,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7649,17 +7644,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_base.py index a90c24e05f..0c184c7f5d 100644 --- a/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/notebook_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,52 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.AssignNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateNotebookExecutionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,52 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.CreateNotebookExecutionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -234,54 +146,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.CreateNotebookRuntimeTemplateRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNotebookExecutionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -292,43 +162,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.DeleteNotebookExecutionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -339,43 +178,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.DeleteNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -386,45 +194,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.DeleteNotebookRuntimeTemplateRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNotebookExecutionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -435,43 +210,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.GetNotebookExecutionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -482,43 +226,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.GetNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -529,43 +242,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.GetNotebookRuntimeTemplateRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNotebookExecutionJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -576,43 +258,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.ListNotebookExecutionJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNotebookRuntimes: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -623,43 +274,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.ListNotebookRuntimesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNotebookRuntimeTemplates: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -670,45 +290,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.ListNotebookRuntimeTemplatesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStartNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -720,52 +307,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.StartNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStopNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -777,38 +324,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.StopNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -817,14 +332,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -836,54 +343,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.UpdateNotebookRuntimeTemplateRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpgradeNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -895,38 +360,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.UpgradeNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -945,17 +378,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -974,17 +396,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1055,17 +466,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1146,22 +546,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1228,17 +612,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1613,17 +986,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2026,17 +1388,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2451,17 +1802,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2872,17 +2212,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3289,16 +2618,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseNotebookServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/persistent_resource_service/client.py b/google/cloud/aiplatform_v1/services/persistent_resource_service/client.py index 89d1a7759f..3d75ba0a35 100644 --- a/google/cloud/aiplatform_v1/services/persistent_resource_service/client.py +++ b/google/cloud/aiplatform_v1/services/persistent_resource_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -145,77 +152,13 @@ class PersistentResourceServiceClient(metaclass=PersistentResourceServiceClientM PersistentResource. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -466,7 +409,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = PersistentResourceServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -487,35 +430,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = PersistentResourceServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -535,67 +455,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = PersistentResourceServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = PersistentResourceServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -729,15 +588,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - PersistentResourceServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( PersistentResourceServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = PersistentResourceServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=PersistentResourceServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -773,14 +634,14 @@ def __init__( self._transport = cast(PersistentResourceServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or PersistentResourceServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=PersistentResourceServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest.py b/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest.py index 03eb4ed48b..1c4a93377d 100644 --- a/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2838,21 +2839,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_create_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource, + "_BaseCreatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2996,17 +2994,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_delete_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource, + "_BaseDeletePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3148,17 +3147,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_get_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource, + "_BaseGetPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3303,17 +3303,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_http_options() ) - request, metadata = self._interceptor.pre_list_persistent_resources( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListPersistentResources, + "_BaseListPersistentResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3463,21 +3464,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_reboot_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource, + "_BaseRebootPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3622,21 +3620,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_update_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource, + "_BaseUpdatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3842,15 +3837,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3986,15 +3982,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4130,15 +4127,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4275,19 +4273,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4424,17 +4419,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4565,17 +4561,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4683,17 +4680,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4804,15 +4802,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4948,15 +4947,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5092,15 +5092,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_asyncio.py index fcfe296d75..925817d871 100644 --- a/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -929,21 +931,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_create_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource, + "_BaseCreatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1097,17 +1096,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource, + "_BaseDeletePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1261,17 +1261,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_get_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource, + "_BaseGetPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1424,17 +1425,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_http_options() ) - request, metadata = await self._interceptor.pre_list_persistent_resources( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListPersistentResources, + "_BaseListPersistentResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1592,21 +1594,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_reboot_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource, + "_BaseRebootPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1761,21 +1760,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_update_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource, + "_BaseUpdatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3972,17 +3968,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4122,17 +4119,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4272,17 +4270,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4423,21 +4422,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4580,17 +4576,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4727,17 +4724,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4849,17 +4847,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4974,17 +4973,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5124,17 +5124,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5274,17 +5275,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_base.py index 734808abb5..251d228357 100644 --- a/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/persistent_resource_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER "persistentResourceId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,54 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.CreatePersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeletePersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -179,45 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.DeletePersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetPersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -228,45 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.GetPersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListPersistentResources: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -277,45 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.ListPersistentResourcesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRebootPersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -327,40 +178,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.RebootPersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdatePersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -369,14 +186,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -388,40 +197,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.UpdatePersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -440,17 +215,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -469,17 +233,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -550,17 +303,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -641,22 +383,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -723,17 +449,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1108,17 +823,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1521,17 +1225,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1946,17 +1639,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2367,17 +2049,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2784,16 +2455,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BasePersistentResourceServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/pipeline_service/client.py b/google/cloud/aiplatform_v1/services/pipeline_service/client.py index c140926a33..3a4f8f8fa1 100644 --- a/google/cloud/aiplatform_v1/services/pipeline_service/client.py +++ b/google/cloud/aiplatform_v1/services/pipeline_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -147,77 +154,13 @@ class PipelineServiceClient(metaclass=PipelineServiceClientMeta): Pipelines). """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -606,7 +549,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = PipelineServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -627,35 +570,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = PipelineServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -675,65 +595,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = PipelineServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = PipelineServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -865,13 +726,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - PipelineServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = PipelineServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = PipelineServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=PipelineServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -907,14 +770,14 @@ def __init__( self._transport = cast(PipelineServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or PipelineServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=PipelineServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=PipelineServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest.py b/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest.py index 7c98c5b063..7a6fe923ff 100644 --- a/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3110,21 +3111,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_http_options() ) - request, metadata = self._interceptor.pre_batch_cancel_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs, + "_BaseBatchCancelPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3269,21 +3267,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_http_options() ) - request, metadata = self._interceptor.pre_batch_delete_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs, + "_BaseBatchDeletePipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3420,21 +3415,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelPipelineJob, + "_BaseCancelPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3537,21 +3529,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline, + "_BaseCancelTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3662,21 +3651,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreatePipelineJob, + "_BaseCreatePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3823,21 +3809,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_create_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline, + "_BaseCreateTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3984,17 +3967,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeletePipelineJob, + "_BaseDeletePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4134,17 +4118,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_delete_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline, + "_BaseDeleteTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4284,17 +4269,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetPipelineJob, + "_BaseGetPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4439,17 +4425,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_get_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetTrainingPipeline, + "_BaseGetTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4591,17 +4578,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListPipelineJobs, + "_BaseListPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4744,17 +4732,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_http_options() ) - request, metadata = self._interceptor.pre_list_training_pipelines( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListTrainingPipelines, + "_BaseListTrainingPipelines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5016,15 +5005,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5157,15 +5147,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5298,15 +5289,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5440,19 +5432,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5587,17 +5576,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5727,17 +5717,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5842,17 +5833,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5960,15 +5952,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6101,15 +6094,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6242,15 +6236,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_asyncio.py index daaa9f8992..f89fb1cc05 100644 --- a/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1229,21 +1231,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_cancel_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs, + "_BaseBatchCancelPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1396,21 +1395,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_delete_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs, + "_BaseBatchDeletePipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1555,21 +1551,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelPipelineJob, + "_BaseCancelPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1678,21 +1671,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline, + "_BaseCancelTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1807,21 +1797,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreatePipelineJob, + "_BaseCreatePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1974,21 +1961,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_create_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline, + "_BaseCreateTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2141,17 +2125,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeletePipelineJob, + "_BaseDeletePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2299,17 +2284,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline, + "_BaseDeleteTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2458,17 +2444,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetPipelineJob, + "_BaseGetPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2621,17 +2608,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_get_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetTrainingPipeline, + "_BaseGetTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2780,17 +2768,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListPipelineJobs, + "_BaseListPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2941,17 +2930,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_http_options() ) - request, metadata = await self._interceptor.pre_list_training_pipelines( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListTrainingPipelines, + "_BaseListTrainingPipelines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5188,17 +5178,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5340,17 +5331,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5492,17 +5484,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5645,21 +5638,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5802,17 +5792,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5949,17 +5940,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6073,17 +6065,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6200,17 +6193,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6352,17 +6346,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6504,17 +6499,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_base.py index e261b6c5d5..a1226a2d0e 100644 --- a/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/pipeline_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,52 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.BatchCancelPipelineJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchDeletePipelineJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -180,52 +132,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.BatchDeletePipelineJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelPipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -237,52 +149,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CancelPipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -294,52 +166,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CancelTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreatePipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -351,52 +183,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CreatePipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -408,52 +200,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CreateTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeletePipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -464,43 +216,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.DeletePipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -511,43 +232,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.DeleteTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetPipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -558,43 +248,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.GetPipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -605,43 +264,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.GetTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListPipelineJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -652,43 +280,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.ListPipelineJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTrainingPipelines: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -699,29 +296,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.ListTrainingPipelinesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -740,17 +314,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -769,17 +332,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -850,17 +402,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -941,22 +482,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1023,17 +548,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1408,17 +922,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1821,17 +1324,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2246,17 +1738,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2667,17 +2148,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3084,16 +2554,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BasePipelineServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/prediction_service/client.py b/google/cloud/aiplatform_v1/services/prediction_service/client.py index a4919f0df2..d33c5f7039 100644 --- a/google/cloud/aiplatform_v1/services/prediction_service/client.py +++ b/google/cloud/aiplatform_v1/services/prediction_service/client.py @@ -41,6 +41,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class PredictionServiceClient(metaclass=PredictionServiceClientMeta): """A service for online predictions and explanations.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -484,7 +427,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = PredictionServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -505,35 +448,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = PredictionServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -553,65 +473,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = PredictionServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = PredictionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -745,13 +606,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - PredictionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = PredictionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = PredictionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=PredictionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -787,14 +650,14 @@ def __init__( self._transport = cast(PredictionServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or PredictionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=PredictionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=PredictionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/prediction_service/transports/rest.py b/google/cloud/aiplatform_v1/services/prediction_service/transports/rest.py index 0c6a9eda75..cc534ca79d 100644 --- a/google/cloud/aiplatform_v1/services/prediction_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/prediction_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -1074,19 +1075,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectPredict._get_http_options() ) - request, metadata = self._interceptor.pre_direct_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectPredict, + "_BaseDirectPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1231,21 +1229,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_http_options() ) - request, metadata = self._interceptor.pre_direct_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectRawPredict, + "_BaseDirectRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1389,19 +1384,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseEmbedContent._get_http_options() ) - request, metadata = self._interceptor.pre_embed_content(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseEmbedContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseEmbedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseEmbedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseEmbedContent, + "_BaseEmbedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1545,23 +1537,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseExplain._get_http_options() ) - request, metadata = self._interceptor.pre_explain(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseExplain._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BaseExplain._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BaseExplain._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseExplain, + "_BaseExplain__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1705,21 +1690,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGenerateContent._get_http_options() ) - request, metadata = self._interceptor.pre_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGenerateContent, + "_BaseGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1863,23 +1845,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BasePredict._get_http_options() ) - request, metadata = self._interceptor.pre_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BasePredict._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BasePredict._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BasePredict._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BasePredict, + "_BasePredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2070,19 +2045,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseRawPredict._get_http_options() ) - request, metadata = self._interceptor.pre_raw_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseRawPredict, + "_BaseRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2231,21 +2203,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_http_options() ) - request, metadata = self._interceptor.pre_server_streaming_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseServerStreamingPredict, + "_BaseServerStreamingPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2422,21 +2391,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_http_options() ) - request, metadata = self._interceptor.pre_stream_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamGenerateContent, + "_BaseStreamGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2661,21 +2627,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_http_options() ) - request, metadata = self._interceptor.pre_stream_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamRawPredict, + "_BaseStreamRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2954,15 +2917,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3096,15 +3060,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3237,15 +3202,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3379,19 +3345,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3526,17 +3489,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3667,17 +3631,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3783,17 +3748,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3901,15 +3867,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4043,15 +4010,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4185,15 +4153,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_asyncio.py index c52e79d525..8c1e0eabf5 100644 --- a/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1197,21 +1199,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_direct_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectPredict, + "_BaseDirectPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1364,21 +1363,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_direct_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectRawPredict, + "_BaseDirectRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1529,21 +1525,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseEmbedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_embed_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseEmbedContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseEmbedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseEmbedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseEmbedContent, + "_BaseEmbedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1695,23 +1688,16 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseExplain._get_http_options() ) - request, metadata = await self._interceptor.pre_explain(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseExplain._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BaseExplain._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BaseExplain._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseExplain, + "_BaseExplain__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1861,21 +1847,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGenerateContent._get_http_options() ) - request, metadata = await self._interceptor.pre_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGenerateContent, + "_BaseGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2025,23 +2008,16 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BasePredict._get_http_options() ) - request, metadata = await self._interceptor.pre_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BasePredict._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BasePredict._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BasePredict._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BasePredict, + "_BasePredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2239,21 +2215,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseRawPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseRawPredict, + "_BaseRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2409,21 +2382,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_server_streaming_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseServerStreamingPredict, + "_BaseServerStreamingPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2605,21 +2575,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_http_options() ) - request, metadata = await self._interceptor.pre_stream_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamGenerateContent, + "_BaseStreamGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2849,21 +2816,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_stream_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamRawPredict, + "_BaseStreamRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3123,17 +3087,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3275,17 +3240,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3427,17 +3393,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3580,21 +3547,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3737,17 +3701,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3884,17 +3849,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4006,17 +3972,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4131,17 +4098,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4283,17 +4251,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4435,17 +4404,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_base.py index c93280a98a..5a073c690a 100644 --- a/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/prediction_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.DirectPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseDirectPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDirectRawPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -176,38 +128,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.DirectRawPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseEmbedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -223,47 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.EmbedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExplain: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -275,52 +160,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.ExplainRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseExplain._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGenerateContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -347,52 +192,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.GenerateContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseGenerateContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -409,52 +214,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.PredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BasePredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRawPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -471,52 +236,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.RawPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseRawPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseServerStreamingPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -533,38 +258,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.StreamingPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamDirectPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -579,14 +272,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -613,38 +298,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.GenerateContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamingPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -659,14 +312,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -683,38 +328,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.StreamRawPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -733,17 +346,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -762,17 +364,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -843,17 +434,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -934,22 +514,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1016,17 +580,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1401,17 +954,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1814,17 +1356,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2239,17 +1770,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2660,17 +2180,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3077,16 +2586,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BasePredictionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/client.py b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/client.py index c5e5be1d0f..5480987c80 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/client.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/client.py @@ -40,6 +40,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -144,77 +151,13 @@ class ReasoningEngineExecutionServiceClient( ): """A service for executing queries on Reasoning Engine.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -402,9 +345,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - ReasoningEngineExecutionServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -425,37 +366,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - ReasoningEngineExecutionServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -475,67 +391,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -669,17 +524,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ReasoningEngineExecutionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( ReasoningEngineExecutionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - ReasoningEngineExecutionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -717,14 +572,14 @@ def __init__( self._transport = cast(ReasoningEngineExecutionServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ReasoningEngineExecutionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest.py b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest.py index 0f86ccaa68..872a1ecd54 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2729,21 +2730,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_async_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine, + "_BaseAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2887,23 +2885,20 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = ( self._interceptor.pre_cancel_async_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine, + "_BaseCancelAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3056,21 +3051,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine, + "_BaseQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3268,21 +3260,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_stream_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine, + "_BaseStreamQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3462,15 +3451,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3606,15 +3596,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3748,15 +3739,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3891,19 +3883,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4040,17 +4029,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4181,17 +4171,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4297,17 +4288,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4416,15 +4408,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4558,15 +4551,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4700,15 +4694,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_asyncio.py index 3db0443f4c..674337f1f5 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -811,23 +813,20 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = ( await self._interceptor.pre_async_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine, + "_BaseAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -981,23 +980,20 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = ( await self._interceptor.pre_cancel_async_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine, + "_BaseCancelAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1158,21 +1154,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine, + "_BaseQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1375,23 +1368,20 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_http_options() ) - request, metadata = ( await self._interceptor.pre_stream_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine, + "_BaseStreamQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3564,17 +3554,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3716,17 +3707,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3868,17 +3860,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4021,21 +4014,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4178,17 +4168,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4327,17 +4318,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4451,17 +4443,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4578,17 +4571,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4730,17 +4724,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4882,17 +4877,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_base.py index 8119e8ef12..ae31174382 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_execution_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -126,56 +118,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - reasoning_engine_execution_service.AsyncQueryReasoningEngineRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelAsyncQueryReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -192,54 +140,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_execution_service.CancelAsyncQueryReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -251,56 +157,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - reasoning_engine_execution_service.QueryReasoningEngineRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamQueryReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -312,42 +174,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - reasoning_engine_execution_service.StreamQueryReasoningEngineRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -366,17 +192,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -395,17 +210,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -476,17 +280,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -567,22 +360,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -649,17 +426,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1034,17 +800,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1447,17 +1202,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1872,17 +1616,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2293,17 +2026,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2710,16 +2432,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseReasoningEngineExecutionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_service/client.py b/google/cloud/aiplatform_v1/services/reasoning_engine_service/client.py index 491b28422c..077ace639f 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_service/client.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -137,77 +144,13 @@ def get_transport_class( class ReasoningEngineServiceClient(metaclass=ReasoningEngineServiceClientMeta): """A service for managing Vertex AI's Reasoning Engines.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -441,7 +384,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ReasoningEngineServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -462,35 +405,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ReasoningEngineServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -510,67 +430,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ReasoningEngineServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ReasoningEngineServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -704,13 +563,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ReasoningEngineServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ReasoningEngineServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ReasoningEngineServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ReasoningEngineServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -746,14 +607,14 @@ def __init__( self._transport = cast(ReasoningEngineServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ReasoningEngineServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ReasoningEngineServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest.py b/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest.py index 176ff788a9..b32409492c 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2775,21 +2776,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_create_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine, + "_BaseCreateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2930,17 +2928,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_delete_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine, + "_BaseDeleteReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3081,17 +3080,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_get_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine, + "_BaseGetReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3236,17 +3236,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_http_options() ) - request, metadata = self._interceptor.pre_list_reasoning_engines( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines, + "_BaseListReasoningEngines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3395,21 +3396,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_update_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine, + "_BaseUpdateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3604,15 +3602,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3746,15 +3745,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3888,15 +3888,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4031,19 +4032,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4178,17 +4176,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4321,17 +4320,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4439,17 +4439,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4560,15 +4561,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4702,15 +4704,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4846,15 +4849,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_asyncio.py index 052a83e5b1..03a10077a5 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -863,21 +865,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_create_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine, + "_BaseCreateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1030,17 +1029,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine, + "_BaseDeleteReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1191,17 +1191,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_get_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine, + "_BaseGetReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1350,17 +1351,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_http_options() ) - request, metadata = await self._interceptor.pre_list_reasoning_engines( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines, + "_BaseListReasoningEngines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1515,21 +1517,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_update_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine, + "_BaseUpdateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3717,17 +3716,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3867,17 +3867,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4017,17 +4018,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4168,21 +4170,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4323,17 +4322,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4470,17 +4470,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4592,17 +4593,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4717,17 +4719,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4867,17 +4870,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5017,17 +5021,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_base.py index 3832355b19..f3226108c6 100644 --- a/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/reasoning_engine_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,54 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.CreateReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,45 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.DeleteReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -226,43 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.GetReasoningEngineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListReasoningEngines: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -273,45 +159,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.ListReasoningEnginesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -323,40 +176,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.UpdateReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -375,17 +194,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -404,17 +212,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -485,17 +282,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -576,22 +362,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -658,17 +428,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1043,17 +802,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1456,17 +1204,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1881,17 +1618,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2302,17 +2028,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2719,16 +2434,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseReasoningEngineServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/schedule_service/client.py b/google/cloud/aiplatform_v1/services/schedule_service/client.py index 7d00d058b2..8f28d6be2b 100644 --- a/google/cloud/aiplatform_v1/services/schedule_service/client.py +++ b/google/cloud/aiplatform_v1/services/schedule_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -141,77 +148,13 @@ class ScheduleServiceClient(metaclass=ScheduleServiceClientMeta): calls. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -644,7 +587,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ScheduleServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -665,35 +608,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ScheduleServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -713,65 +633,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ScheduleServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ScheduleServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -903,13 +764,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ScheduleServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ScheduleServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ScheduleServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ScheduleServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -945,14 +808,14 @@ def __init__( self._transport = cast(ScheduleServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ScheduleServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ScheduleServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/schedule_service/transports/rest.py b/google/cloud/aiplatform_v1/services/schedule_service/transports/rest.py index 4f4c5bdb55..718c4fc591 100644 --- a/google/cloud/aiplatform_v1/services/schedule_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/schedule_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2805,19 +2806,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_create_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCreateSchedule, + "_BaseCreateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2959,15 +2957,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_delete_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteSchedule, + "_BaseDeleteSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3107,15 +3106,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_get_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetSchedule, + "_BaseGetSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3255,15 +3255,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListSchedules._get_http_options() ) - request, metadata = self._interceptor.pre_list_schedules(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListSchedules._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListSchedules._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListSchedules, + "_BaseListSchedules__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3400,19 +3401,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BasePauseSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_pause_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BasePauseSchedule, + "_BasePauseSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3514,19 +3512,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_resume_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseResumeSchedule, + "_BaseResumeSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3636,19 +3631,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_update_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseUpdateSchedule, + "_BaseUpdateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3847,15 +3839,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3988,15 +3981,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4129,15 +4123,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4271,19 +4266,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4418,17 +4410,18 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4558,17 +4551,18 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4673,17 +4667,18 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4791,15 +4786,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4932,15 +4928,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5073,15 +5070,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_asyncio.py index 042c5de8a0..0a393b4078 100644 --- a/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -900,21 +902,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_create_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCreateSchedule, + "_BaseCreateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1065,17 +1064,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteSchedule, + "_BaseDeleteSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1225,17 +1225,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_get_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetSchedule, + "_BaseGetSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1384,17 +1385,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListSchedules._get_http_options() ) - request, metadata = await self._interceptor.pre_list_schedules( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListSchedules._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListSchedules._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListSchedules, + "_BaseListSchedules__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1540,21 +1542,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BasePauseSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_pause_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BasePauseSchedule, + "_BasePauseSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1665,21 +1664,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_resume_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseResumeSchedule, + "_BaseResumeSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1798,21 +1794,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_update_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseUpdateSchedule, + "_BaseUpdateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3998,17 +3991,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4150,17 +4144,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4302,17 +4297,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4455,21 +4451,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4612,17 +4605,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4759,17 +4753,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4883,17 +4878,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5010,17 +5006,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5162,17 +5159,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5314,17 +5312,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_base.py index 4d4a7cdd44..f3316a3975 100644 --- a/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/schedule_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.CreateScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,43 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.DeleteScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,43 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.GetScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseGetSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSchedules: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -271,43 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.ListSchedulesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseListSchedules._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePauseSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -319,52 +178,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.PauseScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BasePauseSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseResumeSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -376,38 +195,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.ResumeScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -416,14 +203,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -435,38 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.UpdateScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -485,17 +232,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -514,17 +250,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -595,17 +320,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -686,22 +400,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -768,17 +466,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1153,17 +840,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1566,17 +1242,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1991,17 +1656,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2412,17 +2066,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2829,16 +2472,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseScheduleServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/session_service/client.py b/google/cloud/aiplatform_v1/services/session_service/client.py index d8c9a49f5a..e3a3da8b91 100644 --- a/google/cloud/aiplatform_v1/services/session_service/client.py +++ b/google/cloud/aiplatform_v1/services/session_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class SessionServiceClient(metaclass=SessionServiceClientMeta): """The service that manages Vertex Session related resources.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -446,7 +389,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = SessionServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -467,35 +410,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = SessionServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -515,65 +435,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = SessionServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = SessionServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = SessionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -703,13 +564,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - SessionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = SessionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = SessionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=SessionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -745,14 +608,14 @@ def __init__( self._transport = cast(SessionServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or SessionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=SessionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=SessionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/session_service/transports/rest.py b/google/cloud/aiplatform_v1/services/session_service/transports/rest.py index 1bd4696100..d20901f022 100644 --- a/google/cloud/aiplatform_v1/services/session_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/session_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2878,19 +2879,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseAppendEvent._get_http_options() ) - request, metadata = self._interceptor.pre_append_event(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseAppendEvent._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseAppendEvent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseAppendEvent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseAppendEvent, + "_BaseAppendEvent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3035,19 +3033,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCreateSession._get_http_options() ) - request, metadata = self._interceptor.pre_create_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseCreateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseCreateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCreateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCreateSession, + "_BaseCreateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3187,15 +3182,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteSession._get_http_options() ) - request, metadata = self._interceptor.pre_delete_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteSession, + "_BaseDeleteSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3333,17 +3329,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetSession._get_http_options() ) - request, metadata = self._interceptor.pre_get_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseGetSession._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetSession, + "_BaseGetSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3483,17 +3478,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListEvents._get_http_options() ) - request, metadata = self._interceptor.pre_list_events(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListEvents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseListEvents._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListEvents, + "_BaseListEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3635,15 +3629,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListSessions._get_http_options() ) - request, metadata = self._interceptor.pre_list_sessions(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListSessions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListSessions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListSessions, + "_BaseListSessions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3786,19 +3781,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseUpdateSession._get_http_options() ) - request, metadata = self._interceptor.pre_update_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseUpdateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseUpdateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseUpdateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseUpdateSession, + "_BaseUpdateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4001,15 +3993,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4142,15 +4135,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4283,15 +4277,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4425,19 +4420,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4571,17 +4563,18 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4711,17 +4704,18 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4826,17 +4820,18 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4944,15 +4939,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5085,15 +5081,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5226,15 +5223,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/session_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/session_service/transports/rest_asyncio.py index e6ceb9815c..ac27a46c7c 100644 --- a/google/cloud/aiplatform_v1/services/session_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/session_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -974,21 +976,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseAppendEvent._get_http_options() ) - request, metadata = await self._interceptor.pre_append_event( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseAppendEvent._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseAppendEvent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseAppendEvent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseAppendEvent, + "_BaseAppendEvent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1141,21 +1140,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCreateSession._get_http_options() ) - request, metadata = await self._interceptor.pre_create_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseCreateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseCreateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCreateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCreateSession, + "_BaseCreateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1305,17 +1301,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteSession._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteSession, + "_BaseDeleteSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1463,19 +1460,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetSession._get_http_options() ) - request, metadata = await self._interceptor.pre_get_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseGetSession._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetSession, + "_BaseGetSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1621,19 +1617,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListEvents._get_http_options() ) - request, metadata = await self._interceptor.pre_list_events( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListEvents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseListEvents._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListEvents, + "_BaseListEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1781,17 +1776,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListSessions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_sessions( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListSessions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListSessions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListSessions, + "_BaseListSessions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1942,21 +1938,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseUpdateSession._get_http_options() ) - request, metadata = await self._interceptor.pre_update_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseUpdateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseUpdateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseUpdateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseUpdateSession, + "_BaseUpdateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4146,17 +4139,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4297,17 +4291,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4448,17 +4443,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4600,21 +4596,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4757,17 +4750,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4904,17 +4898,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5028,17 +5023,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5154,17 +5150,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5306,17 +5303,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5457,17 +5455,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/session_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/session_service/transports/rest_base.py index 78aa285d64..47c2502641 100644 --- a/google/cloud/aiplatform_v1/services/session_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/session_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,52 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.AppendEventRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseAppendEvent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,52 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.CreateSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseCreateSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -233,43 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.DeleteSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseDeleteSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -280,43 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.GetSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseGetSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListEvents: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -327,43 +177,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.ListEventsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseListEvents._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSessions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -374,43 +193,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.ListSessionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseListSessions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -422,38 +210,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.UpdateSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseUpdateSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -472,17 +228,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -501,17 +246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -582,17 +316,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -673,22 +396,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -755,17 +462,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1140,17 +836,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1553,17 +1238,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1978,17 +1652,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2399,17 +2062,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2816,16 +2468,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseSessionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/specialist_pool_service/client.py b/google/cloud/aiplatform_v1/services/specialist_pool_service/client.py index ef0dc4de2b..40bfb5b6ed 100644 --- a/google/cloud/aiplatform_v1/services/specialist_pool_service/client.py +++ b/google/cloud/aiplatform_v1/services/specialist_pool_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -141,77 +148,13 @@ class SpecialistPoolServiceClient(metaclass=SpecialistPoolServiceClientMeta): manage Specialists and tasks on CrowdCompute console. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -399,7 +342,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = SpecialistPoolServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -420,35 +363,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = SpecialistPoolServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -468,67 +388,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = SpecialistPoolServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = SpecialistPoolServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -662,13 +521,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - SpecialistPoolServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = SpecialistPoolServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = SpecialistPoolServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=SpecialistPoolServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -704,14 +565,14 @@ def __init__( self._transport = cast(SpecialistPoolServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or SpecialistPoolServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=SpecialistPoolServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest.py b/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest.py index 2d9b5688d9..933f4a2144 100644 --- a/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2778,21 +2779,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_create_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool, + "_BaseCreateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2935,17 +2933,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_delete_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool, + "_BaseDeleteSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3096,17 +3095,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_get_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool, + "_BaseGetSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3249,17 +3249,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_http_options() ) - request, metadata = self._interceptor.pre_list_specialist_pools( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools, + "_BaseListSpecialistPools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3408,21 +3409,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_update_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool, + "_BaseUpdateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3616,15 +3614,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3758,15 +3757,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3900,15 +3900,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4043,19 +4044,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4190,17 +4188,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4333,17 +4332,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4451,17 +4451,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4572,15 +4573,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4714,15 +4716,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4856,15 +4859,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_asyncio.py index b51892f0c6..a69e7ee34e 100644 --- a/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -862,21 +864,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_create_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool, + "_BaseCreateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1025,17 +1024,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool, + "_BaseDeleteSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1192,17 +1192,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_get_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool, + "_BaseGetSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1349,17 +1350,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_http_options() ) - request, metadata = await self._interceptor.pre_list_specialist_pools( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools, + "_BaseListSpecialistPools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1512,21 +1514,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_update_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool, + "_BaseUpdateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3709,17 +3708,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3859,17 +3859,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4009,17 +4010,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4160,21 +4162,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4315,17 +4314,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4462,17 +4462,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4584,17 +4585,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4709,17 +4711,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4859,17 +4862,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5009,17 +5013,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_base.py index c8f39716d2..ea8bc27e13 100644 --- a/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/specialist_pool_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.CreateSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSpecialistPool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,43 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.DeleteSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetSpecialistPool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -222,43 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.GetSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSpecialistPools: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -269,29 +159,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.ListSpecialistPoolsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateSpecialistPool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -300,14 +167,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -319,38 +178,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.UpdateSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -369,17 +196,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -398,17 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -479,17 +284,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -570,22 +364,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -652,17 +430,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1037,17 +804,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1450,17 +1206,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1875,17 +1620,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2296,17 +2030,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2713,16 +2436,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseSpecialistPoolServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/tensorboard_service/client.py b/google/cloud/aiplatform_v1/services/tensorboard_service/client.py index bc99be3ac8..7a044680d3 100644 --- a/google/cloud/aiplatform_v1/services/tensorboard_service/client.py +++ b/google/cloud/aiplatform_v1/services/tensorboard_service/client.py @@ -40,6 +40,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -149,77 +156,13 @@ def get_transport_class( class TensorboardServiceClient(metaclass=TensorboardServiceClientMeta): """TensorboardService""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -487,7 +430,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = TensorboardServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -508,35 +451,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = TensorboardServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -556,65 +476,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = TensorboardServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = TensorboardServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -748,13 +609,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - TensorboardServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = TensorboardServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = TensorboardServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=TensorboardServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -790,14 +653,14 @@ def __init__( self._transport = cast(TensorboardServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or TensorboardServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=TensorboardServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest.py b/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest.py index 13f9933497..9e364de060 100644 --- a/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -4268,21 +4269,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_tensorboard_runs( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns, + "_BaseBatchCreateTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4434,23 +4432,20 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_create_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries, + "_BaseBatchCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4601,19 +4596,20 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_read_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData, + "_BaseBatchReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4762,21 +4758,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboard, + "_BaseCreateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4920,21 +4913,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment, + "_BaseCreateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5085,21 +5075,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun, + "_BaseCreateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5247,21 +5234,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries, + "_BaseCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5410,17 +5394,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard, + "_BaseDeleteTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5561,17 +5546,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment, + "_BaseDeleteTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5713,17 +5699,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun, + "_BaseDeleteTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5866,17 +5853,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries, + "_BaseDeleteTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6021,23 +6009,20 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( self._interceptor.pre_export_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData, + "_BaseExportTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6189,15 +6174,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboard, + "_BaseGetTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6341,17 +6327,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment, + "_BaseGetTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6498,17 +6485,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun, + "_BaseGetTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6650,17 +6638,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries, + "_BaseGetTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6806,17 +6795,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboard_experiments( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments, + "_BaseListTensorboardExperiments__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6961,17 +6951,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboard_runs( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns, + "_BaseListTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7118,17 +7109,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboards( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboards, + "_BaseListTensorboards__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7272,17 +7264,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries, + "_BaseListTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7429,17 +7422,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_blob_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData, + "_BaseReadTensorboardBlobData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7576,17 +7570,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_size( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize, + "_BaseReadTensorboardSize__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7734,17 +7729,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_time_series_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData, + "_BaseReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7889,17 +7885,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_usage( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage, + "_BaseReadTensorboardUsage__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8048,21 +8045,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard, + "_BaseUpdateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8206,21 +8200,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment, + "_BaseUpdateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8371,21 +8362,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun, + "_BaseUpdateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8533,21 +8521,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries, + "_BaseUpdateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8699,21 +8684,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_http_options() ) - request, metadata = self._interceptor.pre_write_tensorboard_experiment_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData, + "_BaseWriteTensorboardExperimentData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8863,21 +8845,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_http_options() ) - request, metadata = self._interceptor.pre_write_tensorboard_run_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData, + "_BaseWriteTensorboardRunData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9347,15 +9326,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9489,15 +9469,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9631,15 +9612,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9774,19 +9756,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9921,17 +9900,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10064,17 +10044,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10180,17 +10161,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10299,15 +10281,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10441,15 +10424,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10583,15 +10567,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_asyncio.py index 90d08cb290..7a6a2844fc 100644 --- a/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -2480,23 +2482,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_create_tensorboard_runs( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns, + "_BaseBatchCreateTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2654,23 +2653,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_create_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries, + "_BaseBatchCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2829,19 +2825,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_read_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData, + "_BaseBatchReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2998,21 +2995,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_create_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboard, + "_BaseCreateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3166,23 +3160,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_tensorboard_experiment( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment, + "_BaseCreateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3339,21 +3330,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_create_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun, + "_BaseCreateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3507,23 +3495,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries, + "_BaseCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3678,17 +3663,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard, + "_BaseDeleteTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3839,19 +3825,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_tensorboard_experiment( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment, + "_BaseDeleteTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4001,17 +3988,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun, + "_BaseDeleteTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4162,19 +4150,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries, + "_BaseDeleteTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4327,23 +4316,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_export_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData, + "_BaseExportTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4503,17 +4489,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboard, + "_BaseGetTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4663,17 +4650,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment, + "_BaseGetTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4826,17 +4814,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun, + "_BaseGetTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4984,17 +4973,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries, + "_BaseGetTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5148,19 +5138,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_tensorboard_experiments( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments, + "_BaseListTensorboardExperiments__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5313,17 +5304,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_http_options() ) - request, metadata = await self._interceptor.pre_list_tensorboard_runs( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns, + "_BaseListTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5474,17 +5466,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_http_options() ) - request, metadata = await self._interceptor.pre_list_tensorboards( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboards, + "_BaseListTensorboards__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5636,19 +5629,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries, + "_BaseListTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5802,17 +5796,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_http_options() ) - request, metadata = await self._interceptor.pre_read_tensorboard_blob_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData, + "_BaseReadTensorboardBlobData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5955,17 +5950,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_http_options() ) - request, metadata = await self._interceptor.pre_read_tensorboard_size( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize, + "_BaseReadTensorboardSize__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6119,19 +6115,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_read_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData, + "_BaseReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6282,17 +6279,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_http_options() ) - request, metadata = await self._interceptor.pre_read_tensorboard_usage( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage, + "_BaseReadTensorboardUsage__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6445,21 +6443,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_update_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard, + "_BaseUpdateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6613,23 +6608,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_tensorboard_experiment( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment, + "_BaseUpdateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6786,21 +6778,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_update_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun, + "_BaseUpdateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6954,23 +6943,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries, + "_BaseUpdateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7128,23 +7114,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_write_tensorboard_experiment_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData, + "_BaseWriteTensorboardExperimentData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7300,21 +7283,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_http_options() ) - request, metadata = await self._interceptor.pre_write_tensorboard_run_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData, + "_BaseWriteTensorboardRunData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9723,17 +9703,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9875,17 +9856,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10027,17 +10009,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10180,21 +10163,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10337,17 +10317,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10484,17 +10465,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10606,17 +10588,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10731,17 +10714,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10883,17 +10867,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11033,17 +11018,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_base.py index df046c1782..32aeb8c86a 100644 --- a/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/tensorboard_service/transports/rest_base.py @@ -110,14 +110,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -129,54 +121,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.BatchCreateTensorboardRunsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchCreateTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -188,40 +138,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.BatchCreateTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchReadTensorboardTimeSeriesData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -230,14 +146,6 @@ def __hash__(self): # pragma: NO COVER "timeSeries": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -248,47 +156,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - tensorboard_service.BatchReadTensorboardTimeSeriesDataRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -300,38 +173,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -340,14 +181,6 @@ def __hash__(self): # pragma: NO COVER "tensorboardExperimentId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -359,40 +192,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardExperimentRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -401,14 +200,6 @@ def __hash__(self): # pragma: NO COVER "tensorboardRunId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -420,52 +211,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -477,54 +228,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -535,43 +244,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -582,45 +260,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardExperimentRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -631,43 +276,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -678,45 +292,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportTensorboardTimeSeriesData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -728,54 +309,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ExportTensorboardTimeSeriesDataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -786,43 +325,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -833,43 +341,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardExperimentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -880,43 +357,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -927,43 +373,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardTimeSeriesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboardExperiments: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -974,45 +389,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardExperimentsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboardRuns: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1023,43 +405,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardRunsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboards: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1070,43 +421,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1117,45 +437,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardBlobData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1166,43 +453,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardBlobDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardSize: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1213,43 +469,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardSizeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardTimeSeriesData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1260,45 +485,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardTimeSeriesDataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardUsage: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1309,29 +501,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardUsageRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1340,14 +509,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1359,38 +520,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1399,14 +528,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1418,40 +539,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardExperimentRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1460,14 +547,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1479,38 +558,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1519,14 +566,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1538,54 +577,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseWriteTensorboardExperimentData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1597,54 +594,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.WriteTensorboardExperimentDataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseWriteTensorboardRunData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1656,38 +611,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.WriteTensorboardRunDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1706,17 +629,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1735,17 +647,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1816,17 +717,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1907,22 +797,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1989,17 +863,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2374,17 +1237,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2787,17 +1639,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3212,17 +2053,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3633,17 +2463,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4050,16 +2869,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseTensorboardServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/client.py b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/client.py index 32d40b6f54..2fd065a00e 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/client.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -137,77 +144,13 @@ def get_transport_class( class VertexRagDataServiceClient(metaclass=VertexRagDataServiceClientMeta): """A service for managing user data for RAG.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -505,7 +448,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = VertexRagDataServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -526,35 +469,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = VertexRagDataServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -574,65 +494,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = VertexRagDataServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = VertexRagDataServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -766,13 +627,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - VertexRagDataServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = VertexRagDataServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = VertexRagDataServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=VertexRagDataServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -808,14 +671,14 @@ def __init__( self._transport = cast(VertexRagDataServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or VertexRagDataServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=VertexRagDataServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest.py b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest.py index 34a5ed7fd5..668ca53578 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3180,21 +3181,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_create_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus, + "_BaseCreateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3335,17 +3333,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_delete_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus, + "_BaseDeleteRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3485,15 +3484,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_http_options() ) - request, metadata = self._interceptor.pre_delete_rag_file(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile, + "_BaseDeleteRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3633,15 +3633,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_corpus(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus, + "_BaseGetRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3780,17 +3781,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig, + "_BaseGetRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3933,15 +3935,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_file(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagFile, + "_BaseGetRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4084,21 +4087,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_http_options() ) - request, metadata = self._interceptor.pre_import_rag_files( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles, + "_BaseImportRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4238,17 +4238,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_http_options() ) - request, metadata = self._interceptor.pre_list_rag_corpora( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora, + "_BaseListRagCorpora__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4391,15 +4392,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_http_options() ) - request, metadata = self._interceptor.pre_list_rag_files(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagFiles, + "_BaseListRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4544,21 +4546,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_update_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus, + "_BaseUpdateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4700,21 +4699,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_http_options() ) - request, metadata = self._interceptor.pre_update_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig, + "_BaseUpdateRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4857,19 +4853,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_http_options() ) - request, metadata = self._interceptor.pre_upload_rag_file(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile, + "_BaseUploadRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5135,15 +5128,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5277,15 +5271,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5419,15 +5414,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5562,19 +5558,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5709,17 +5702,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5852,17 +5846,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5968,17 +5963,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6087,15 +6083,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6229,15 +6226,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6371,15 +6369,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_asyncio.py index f931d5007c..1a9c62808c 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1299,21 +1301,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_create_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus, + "_BaseCreateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1462,17 +1461,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus, + "_BaseDeleteRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1620,17 +1620,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_rag_file( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile, + "_BaseDeleteRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1778,17 +1779,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus, + "_BaseGetRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1933,17 +1935,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig, + "_BaseGetRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2090,17 +2093,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_file( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagFile, + "_BaseGetRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2251,21 +2255,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_http_options() ) - request, metadata = await self._interceptor.pre_import_rag_files( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles, + "_BaseImportRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2413,17 +2414,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_http_options() ) - request, metadata = await self._interceptor.pre_list_rag_corpora( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora, + "_BaseListRagCorpora__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2572,17 +2574,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_http_options() ) - request, metadata = await self._interceptor.pre_list_rag_files( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagFiles, + "_BaseListRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2733,21 +2736,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_update_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus, + "_BaseUpdateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2897,21 +2897,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_http_options() ) - request, metadata = await self._interceptor.pre_update_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig, + "_BaseUpdateRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3062,21 +3059,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_http_options() ) - request, metadata = await self._interceptor.pre_upload_rag_file( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile, + "_BaseUploadRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5317,17 +5311,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5469,17 +5464,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5619,17 +5615,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5770,21 +5767,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5925,17 +5919,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6072,17 +6067,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6194,17 +6190,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6319,17 +6316,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6469,17 +6467,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6619,17 +6618,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_base.py index ab37fc8817..d2dc327b44 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_data_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.CreateRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteRagCorpus: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,43 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.DeleteRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteRagFile: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -222,43 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.DeleteRagFileRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagCorpus: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -269,43 +159,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagEngineConfig: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -316,43 +175,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagEngineConfigRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagFile: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -363,43 +191,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagFileRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportRagFiles: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -411,52 +208,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ImportRagFilesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListRagCorpora: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -467,43 +224,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ListRagCorporaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListRagFiles: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -514,43 +240,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ListRagFilesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateRagCorpus: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -562,52 +257,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.UpdateRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateRagEngineConfig: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -619,54 +274,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.UpdateRagEngineConfigRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUploadRagFile: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -678,38 +291,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.UploadRagFileRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -728,17 +309,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -757,17 +327,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -838,17 +397,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -929,22 +477,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1011,17 +543,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1396,17 +917,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1809,17 +1319,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2234,17 +1733,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2655,17 +2143,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3072,16 +2549,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseVertexRagDataServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_service/client.py b/google/cloud/aiplatform_v1/services/vertex_rag_service/client.py index dae9a3a34e..62f3f0c626 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_service/client.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -132,77 +139,13 @@ def get_transport_class( class VertexRagServiceClient(metaclass=VertexRagServiceClientMeta): """A service for retrieving relevant contexts.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -390,7 +333,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = VertexRagServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -411,35 +354,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = VertexRagServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -459,65 +379,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = VertexRagServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = VertexRagServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -649,13 +510,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - VertexRagServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = VertexRagServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = VertexRagServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=VertexRagServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -691,14 +554,14 @@ def __init__( self._transport = cast(VertexRagServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or VertexRagServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=VertexRagServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest.py b/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest.py index 761e07adcf..2d4b7360f8 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2776,19 +2777,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAskContexts._get_http_options() ) - request, metadata = self._interceptor.pre_ask_contexts(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAskContexts, + "_BaseAskContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2934,21 +2932,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_http_options() ) - request, metadata = self._interceptor.pre_async_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts, + "_BaseAsyncRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3087,19 +3082,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_http_options() ) - request, metadata = self._interceptor.pre_augment_prompt(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAugmentPrompt, + "_BaseAugmentPrompt__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3244,21 +3236,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_http_options() ) - request, metadata = self._interceptor.pre_corroborate_content( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCorroborateContent, + "_BaseCorroborateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3403,21 +3392,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_http_options() ) - request, metadata = self._interceptor.pre_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseRetrieveContexts, + "_BaseRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3613,15 +3599,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3754,15 +3741,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3895,15 +3883,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4037,19 +4026,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4184,17 +4170,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4325,17 +4312,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4441,17 +4429,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4559,15 +4548,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4700,15 +4690,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4841,15 +4832,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_asyncio.py index e82299ca17..a3673f69ca 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -861,21 +863,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAskContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_ask_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAskContexts, + "_BaseAskContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1029,21 +1028,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_async_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts, + "_BaseAsyncRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1191,21 +1187,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_http_options() ) - request, metadata = await self._interceptor.pre_augment_prompt( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAugmentPrompt, + "_BaseAugmentPrompt__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1358,21 +1351,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_http_options() ) - request, metadata = await self._interceptor.pre_corroborate_content( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCorroborateContent, + "_BaseCorroborateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1523,21 +1513,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseRetrieveContexts, + "_BaseRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3723,17 +3710,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3875,17 +3863,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4027,17 +4016,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4180,21 +4170,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4337,17 +4324,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4484,17 +4472,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4608,17 +4597,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4735,17 +4725,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4887,17 +4878,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5039,17 +5031,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_base.py index 4c834c5bfb..176f65ceea 100644 --- a/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/vertex_rag_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,52 +110,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.AskContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseAskContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAsyncRetrieveContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,52 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.AsyncRetrieveContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAugmentPrompt: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -232,52 +144,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.AugmentPromptRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCorroborateContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -289,52 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.CorroborateContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRetrieveContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -346,38 +178,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.RetrieveContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -396,17 +196,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -425,17 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -506,17 +284,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -597,22 +364,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -679,17 +430,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1064,17 +804,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1477,17 +1206,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1902,17 +1620,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2323,17 +2030,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2740,16 +2436,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseVertexRagServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1/services/vizier_service/client.py b/google/cloud/aiplatform_v1/services/vizier_service/client.py index 1fec5d486a..8164124aa0 100644 --- a/google/cloud/aiplatform_v1/services/vizier_service/client.py +++ b/google/cloud/aiplatform_v1/services/vizier_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -136,77 +143,13 @@ class VizierServiceClient(metaclass=VizierServiceClientMeta): searching over deep learning architectures. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -440,7 +383,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = VizierServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -461,35 +404,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = VizierServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -509,65 +429,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = VizierServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = VizierServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = VizierServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -697,13 +558,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - VizierServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = VizierServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = VizierServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=VizierServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -739,14 +602,14 @@ def __init__( self._transport = cast(VizierServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or VizierServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=VizierServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=VizierServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1/services/vizier_service/transports/rest.py b/google/cloud/aiplatform_v1/services/vizier_service/transports/rest.py index 90030fdc05..0226e18de7 100644 --- a/google/cloud/aiplatform_v1/services/vizier_service/transports/rest.py +++ b/google/cloud/aiplatform_v1/services/vizier_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3231,21 +3232,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_http_options() ) - request, metadata = self._interceptor.pre_add_trial_measurement( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseAddTrialMeasurement, + "_BaseAddTrialMeasurement__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3390,21 +3388,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_http_options() ) - request, metadata = self._interceptor.pre_check_trial_early_stopping_state( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState, + "_BaseCheckTrialEarlyStoppingState__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3551,19 +3546,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCompleteTrial._get_http_options() ) - request, metadata = self._interceptor.pre_complete_trial(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCompleteTrial, + "_BaseCompleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3703,23 +3695,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateStudy._get_http_options() ) - request, metadata = self._interceptor.pre_create_study(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateStudy, + "_BaseCreateStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3864,23 +3849,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateTrial._get_http_options() ) - request, metadata = self._interceptor.pre_create_trial(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateTrial._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateTrial, + "_BaseCreateTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4015,17 +3993,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteStudy._get_http_options() ) - request, metadata = self._interceptor.pre_delete_study(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteStudy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteStudy, + "_BaseDeleteStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4125,17 +4102,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteTrial._get_http_options() ) - request, metadata = self._interceptor.pre_delete_trial(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteTrial._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteTrial, + "_BaseDeleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4239,19 +4215,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetStudy._get_http_options() ) - request, metadata = self._interceptor.pre_get_study(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetStudy, + "_BaseGetStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4394,19 +4367,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetTrial._get_http_options() ) - request, metadata = self._interceptor.pre_get_trial(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetTrial, + "_BaseGetTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4547,21 +4517,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_http_options() ) - request, metadata = self._interceptor.pre_list_optimal_trials( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOptimalTrials, + "_BaseListOptimalTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4704,17 +4671,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListStudies._get_http_options() ) - request, metadata = self._interceptor.pre_list_studies(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseListStudies._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListStudies._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListStudies, + "_BaseListStudies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4856,19 +4822,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListTrials._get_http_options() ) - request, metadata = self._interceptor.pre_list_trials(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListTrials, + "_BaseListTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5009,23 +4972,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseLookupStudy._get_http_options() ) - request, metadata = self._interceptor.pre_lookup_study(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseLookupStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseLookupStudy, + "_BaseLookupStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5170,25 +5126,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseStopTrial._get_http_options() ) - request, metadata = self._interceptor.pre_stop_trial(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseStopTrial, + "_BaseStopTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5331,19 +5278,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSuggestTrials._get_http_options() ) - request, metadata = self._interceptor.pre_suggest_trials(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSuggestTrials, + "_BaseSuggestTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5605,17 +5549,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5748,15 +5691,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5889,15 +5833,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6031,19 +5976,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6177,17 +6119,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6317,17 +6260,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6432,17 +6376,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6550,15 +6495,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6691,15 +6637,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6832,15 +6779,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_asyncio.py index e4e14c7031..68ff363c43 100644 --- a/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1366,21 +1368,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_http_options() ) - request, metadata = await self._interceptor.pre_add_trial_measurement( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseAddTrialMeasurement, + "_BaseAddTrialMeasurement__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1531,23 +1530,20 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_http_options() ) - request, metadata = ( await self._interceptor.pre_check_trial_early_stopping_state( request, metadata ) ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState, + "_BaseCheckTrialEarlyStoppingState__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1700,21 +1696,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCompleteTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_complete_trial( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCompleteTrial, + "_BaseCompleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1862,25 +1855,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_create_study( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateStudy, + "_BaseCreateStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2031,25 +2017,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_create_trial( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateTrial._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateTrial, + "_BaseCreateTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2190,19 +2169,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_study( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteStudy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteStudy, + "_BaseDeleteStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2308,19 +2286,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_trial( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteTrial._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteTrial, + "_BaseDeleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2430,19 +2407,16 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_study(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetStudy, + "_BaseGetStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2591,19 +2565,16 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_get_trial(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetTrial, + "_BaseGetTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2751,21 +2722,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_http_options() ) - request, metadata = await self._interceptor.pre_list_optimal_trials( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOptimalTrials, + "_BaseListOptimalTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2916,19 +2884,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListStudies._get_http_options() ) - request, metadata = await self._interceptor.pre_list_studies( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListStudies._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListStudies._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListStudies, + "_BaseListStudies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3076,21 +3043,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListTrials._get_http_options() ) - request, metadata = await self._interceptor.pre_list_trials( request, metadata ) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListTrials, + "_BaseListTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3237,25 +3201,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseLookupStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_lookup_study( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseLookupStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseLookupStudy, + "_BaseLookupStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3406,27 +3363,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseStopTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_stop_trial( request, metadata ) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseStopTrial, + "_BaseStopTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3575,21 +3523,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSuggestTrials._get_http_options() ) - request, metadata = await self._interceptor.pre_suggest_trials( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSuggestTrials, + "_BaseSuggestTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5824,19 +5769,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5975,17 +5919,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6126,17 +6071,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6278,21 +6224,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6435,17 +6378,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6583,17 +6527,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6706,17 +6651,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6832,17 +6778,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6983,17 +6930,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7134,17 +7082,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_base.py b/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_base.py index 988b5466f3..61ad1573d7 100644 --- a/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1/services/vizier_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.AddTrialMeasurementRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCheckTrialEarlyStoppingState: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -178,52 +130,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CheckTrialEarlyStoppingStateRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCompleteTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -235,52 +147,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CompleteTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCompleteTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -292,52 +164,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CreateStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -349,52 +181,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CreateTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -405,43 +197,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.DeleteStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseDeleteStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -452,43 +213,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.DeleteTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseDeleteTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -499,43 +229,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.GetStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseGetStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -546,43 +245,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.GetTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseGetTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListOptimalTrials: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -594,52 +262,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.ListOptimalTrialsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListStudies: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -650,43 +278,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.ListStudiesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseListStudies._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTrials: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -697,43 +294,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.ListTrialsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseListTrials._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseLookupStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -745,52 +311,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.LookupStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStopTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -802,52 +328,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.StopTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseStopTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSuggestTrials: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -859,38 +345,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.SuggestTrialsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseSuggestTrials._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -909,17 +363,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -938,17 +381,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1019,17 +451,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1110,22 +531,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1192,17 +597,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1577,17 +971,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1990,17 +1373,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2415,17 +1787,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2836,17 +2197,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3253,16 +2603,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseVizierServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/__init__.py b/google/cloud/aiplatform_v1beta1/__init__.py index 3ea145077f..a15f7d84a8 100644 --- a/google/cloud/aiplatform_v1beta1/__init__.py +++ b/google/cloud/aiplatform_v1beta1/__init__.py @@ -16,12 +16,9 @@ from google.cloud.aiplatform_v1beta1 import gapic_version as package_version import google.api_core as api_core -import sys __version__ = package_version.__version__ -from importlib import metadata - # PEP 0810: Explicit Lazy Imports # Python 3.15+ natively intercepts and defers these imports. # Developers can disable this behavior and force eager imports. @@ -984,7 +981,18 @@ from .types.match_service import FindNeighborsResponse from .types.match_service import ReadIndexDatapointsRequest from .types.match_service import ReadIndexDatapointsResponse +from .types.memory_bank import IntermediateExtractedMemory from .types.memory_bank import Memory +from .types.memory_bank import MemoryBankCustomizationConfig +from .types.memory_bank import MemoryConjunctionFilter +from .types.memory_bank import MemoryFilter +from .types.memory_bank import MemoryGenerationTriggerConfig +from .types.memory_bank import MemoryMetadataValue +from .types.memory_bank import MemoryProfile +from .types.memory_bank import MemoryRevision +from .types.memory_bank import MemoryTopicId +from .types.memory_bank import StructuredMemoryConfig +from .types.memory_bank import MemoryType from .types.memory_bank_service import CreateMemoryOperationMetadata from .types.memory_bank_service import CreateMemoryRequest from .types.memory_bank_service import DeleteMemoryOperationMetadata @@ -1616,89 +1624,6 @@ from .types.vizier_service import SuggestTrialsRequest from .types.vizier_service import SuggestTrialsResponse -if hasattr(api_core, "check_python_version") and hasattr( - api_core, "check_dependency_versions" -): # pragma: NO COVER - api_core.check_python_version("google.cloud.aiplatform_v1beta1") # type: ignore - api_core.check_dependency_versions("google.cloud.aiplatform_v1beta1") # type: ignore -else: # pragma: NO COVER - # An older version of api_core is installed which does not define the - # functions above. We do equivalent checks manually. - try: - import warnings - - _py_version_str = sys.version.split()[0] - _package_label = "google.cloud.aiplatform_v1beta1" - if sys.version_info < (3, 10): - warnings.warn( - "You are using a non-supported Python version " - + f"({_py_version_str}). Google will not post any further " - + f"updates to {_package_label} supporting this Python version. " - + "Please upgrade to the latest Python version, or at " - + f"least to Python 3.10, and then update {_package_label}.", - FutureWarning, - ) - - def parse_version_to_tuple(version_string: str): - """Safely converts a semantic version string to a comparable tuple of integers. - Example: "6.33.5" -> (6, 33, 5) - Ignores non-numeric parts and handles common version formats. - Args: - version_string: Version string in the format "x.y.z" or "x.y.z" - Returns: - Tuple of integers for the parsed version string. - """ - parts = [] - for part in version_string.split("."): - try: - parts.append(int(part)) - except ValueError: - # If it's a non-numeric part (e.g., '1.0.0b1' -> 'b1'), stop here. - # This is a simplification compared to 'packaging.parse_version', but sufficient - # for comparing strictly numeric semantic versions. - break - return tuple(parts) - - def _get_version(dependency_name): - try: - version_string: str = metadata.version(dependency_name) - parsed_version = parse_version_to_tuple(version_string) - return (parsed_version, version_string) - except Exception: - # Catch exceptions from metadata.version() (e.g., PackageNotFoundError) - # or errors during parse_version_to_tuple - return (None, "--") - - _dependency_package = "google.protobuf" - _next_supported_version = "6.33.5" - _next_supported_version_tuple = (6, 33, 5) - _recommendation = " (we recommend 7.x)" - (_version_used, _version_used_string) = _get_version(_dependency_package) - if _version_used and _version_used < _next_supported_version_tuple: - warnings.warn( - f"Package {_package_label} depends on " - + f"{_dependency_package}, currently installed at version " - + f"{_version_used_string}. Future updates to " - + f"{_package_label} will require {_dependency_package} at " - + f"version {_next_supported_version} or higher{_recommendation}." - + " Please ensure " - + "that either (a) your Python environment doesn't pin the " - + f"version of {_dependency_package}, so that updates to " - + f"{_package_label} can require the higher version, or " - + "(b) you manually update your Python environment to use at " - + f"least version {_next_supported_version} of " - + f"{_dependency_package}.", - FutureWarning, - ) - except Exception: - warnings.warn( - "Could not determine the version of Python " - + "currently being used. To continue receiving " - + "updates for {_package_label}, ensure you are " - + "using a supported version of Python; see " - + "https://devguide.python.org/versions/" - ) - __all__ = ( "DatasetServiceAsyncClient", "DeploymentResourcePoolServiceAsyncClient", @@ -2329,6 +2254,7 @@ def _get_version(dependency_name): "InputDataConfig", "Int64Array", "IntegratedGradientsAttribution", + "IntermediateExtractedMemory", "JiraSource", "JobServiceClient", "JobState", @@ -2478,7 +2404,16 @@ def _get_version(dependency_name): "MatchServiceClient", "Measurement", "Memory", + "MemoryBankCustomizationConfig", "MemoryBankServiceClient", + "MemoryConjunctionFilter", + "MemoryFilter", + "MemoryGenerationTriggerConfig", + "MemoryMetadataValue", + "MemoryProfile", + "MemoryRevision", + "MemoryTopicId", + "MemoryType", "MergeVersionAliasesRequest", "MetadataList", "MetadataSchema", @@ -2827,6 +2762,7 @@ def _get_version(dependency_name): "StringArray", "StructFieldValue", "StructValue", + "StructuredMemoryConfig", "Study", "StudySpec", "StudyTimeConstraint", @@ -3045,3 +2981,6 @@ def _get_version(dependency_name): "WriteTensorboardRunDataResponse", "XraiAttribution", ) + +api_core.check_python_version("google.cloud.aiplatform_v1beta1") +api_core.check_dependency_versions("google.cloud.aiplatform_v1beta1") diff --git a/google/cloud/aiplatform_v1beta1/_compat.py b/google/cloud/aiplatform_v1beta1/_compat.py new file mode 100644 index 0000000000..9152466482 --- /dev/null +++ b/google/cloud/aiplatform_v1beta1/_compat.py @@ -0,0 +1,263 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""A compatibility module for older versions of google-api-core.""" + +import os +import json + +from typing import Any, Dict, List, Optional, Tuple + +from google.api_core import path_template +from google.api_core.universe import EmptyUniverseError +from google.auth.exceptions import MutualTLSChannelError +from google.protobuf import json_format +from urllib.parse import urlparse, urlunparse + +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 + # we bump the minimum supported version of google-auth. + from google.auth.transport.mtls import should_use_client_cert # type: ignore +except ImportError: # pragma: NO COVER + + def should_use_client_cert(): + """Returns whether client certificate should be used for mTLS.""" + use_client_cert = os.getenv( + "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" + ).lower() + if use_client_cert not in ("true", "false"): + raise ValueError( + "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" + " either `true` or `false`" + ) + return use_client_cert == "true" + + +def read_environment_variables(): + """Returns the environment variables used by the client. + + Returns: + Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, + GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. + + Raises: + ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not + any of ["true", "false"]. + google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT + is not any of ["auto", "never", "always"]. + """ + use_client_cert = should_use_client_cert() + use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() + universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") + if use_mtls_endpoint not in ("auto", "never", "always"): + 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 + + +DEFAULT_UNIVERSE = "googleapis.com" + +try: + from google.api_core.universe import get_default_mtls_endpoint +except ImportError: # pragma: NO COVER + + def get_default_mtls_endpoint(api_endpoint: Optional[str]) -> Optional[str]: + """Converts api endpoint to mTLS endpoint. + + Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to + "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. + Other URLs (including those that do not match these domain suffixes or + already contain '.mtls.') are passed through as-is. + + Args: + api_endpoint (Optional[str]): the api endpoint to convert. + + Returns: + Optional[str]: converted mTLS api endpoint. + """ + if not api_endpoint or ".mtls." in api_endpoint.lower(): + return api_endpoint + + has_scheme = "://" in api_endpoint + if not has_scheme: + parsed = urlparse("//" + api_endpoint) + else: + parsed = urlparse(api_endpoint) + + host = parsed.hostname + if not host: + return api_endpoint + + port = f":{parsed.port}" if parsed.port else "" + + lowered_host = host.lower() + suffix_sandbox = ".sandbox.googleapis.com" + suffix_google = ".googleapis.com" + if lowered_host.endswith(suffix_sandbox): + new_host = host[: -len(suffix_sandbox)] + ".mtls.sandbox.googleapis.com" + elif lowered_host.endswith(suffix_google): + new_host = host[: -len(suffix_google)] + ".mtls.googleapis.com" + else: + return api_endpoint + + netloc = new_host + port + new_parsed = parsed._replace(netloc=netloc) + + if not has_scheme: + return urlunparse(new_parsed)[2:] + else: + return urlunparse(new_parsed) + + +try: + from google.api_core.universe import get_api_endpoint +except ImportError: # pragma: NO COVER + + def get_api_endpoint( + api_override: Optional[str], + universe_domain: str, + default_universe: str, + default_mtls_endpoint: Optional[str], + default_endpoint_template: str, + use_mtls: bool, + ) -> str: + """Return the API endpoint used by the client. + + Args: + api_override (Optional[str]): The API endpoint override. If specified, + this is always returned. + universe_domain (str): The universe domain used by the client. + default_universe (str): The default universe domain. + default_mtls_endpoint (Optional[str]): The default mTLS endpoint. + default_endpoint_template (str): The default endpoint template containing + a placeholder `{UNIVERSE_DOMAIN}`. + use_mtls (bool): Whether to use the mTLS endpoint. + + Returns: + str: The API endpoint to be used by the client. + + Raises: + google.auth.exceptions.MutualTLSChannelError: If mTLS is requested but + not supported in the configured universe domain. + ValueError: If mTLS is requested but no mTLS endpoint is available. + """ + if api_override is not None: + return api_override + + if use_mtls: + if universe_domain.lower() != default_universe.lower(): + raise MutualTLSChannelError( + f"mTLS is not supported in any universe other than {default_universe}." + ) + if not default_mtls_endpoint: + raise ValueError("mTLS endpoint is not available.") + return default_mtls_endpoint + else: + return default_endpoint_template.format(UNIVERSE_DOMAIN=universe_domain) + + +try: + from google.api_core.universe import get_universe_domain +except ImportError: # pragma: NO COVER + + def get_universe_domain( + *potential_universes: Optional[str], + default_universe: str, + ) -> str: + """Return the universe domain used by the client. + + Args: + *potential_universes (Optional[str]): Potential universe domains in order of preference. + default_universe (str): The default universe domain. + + Returns: + str: The universe domain to be used by the client. + + Raises: + EmptyUniverseError: If the resolved universe domain is an empty string. + """ + resolved = next( + (x.strip() for x in potential_universes if x is not None), + default_universe, + ) + + if not resolved: + raise EmptyUniverseError() + return resolved + + +try: + from google.api_core.rest_helpers import transcode_request # type: ignore +except ImportError: # pragma: NO COVER + + def transcode_request( + http_options: List[Dict[str, str]], + request: Any, + required_fields_default_values: Optional[Dict[str, Any]] = None, + rest_numeric_enums: bool = False, + ) -> Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: + """Transcodes a request into HTTP method, URI, body, and query parameters. + + Args: + http_options (List[Dict[str, str]]): List of HTTP transcoding rules. + request (Any): The protobuf or proto-plus request message. + required_fields_default_values (Optional[Dict[str, Any]]): Dictionary + of required fields default values to merge into query parameters if missing. + rest_numeric_enums (bool): Whether to encode enums as integers. + + Returns: + Tuple[Dict[str, Any], Optional[str], Dict[str, Any]]: A tuple containing: + - The raw transcoded request dictionary (containing keys like 'uri', 'method'). + - The serialized request body JSON string, or None if no body. + - The query parameters dictionary. + """ + if request is None: + raise TypeError("request cannot be None") + + # Convert proto-plus message to its underlying protobuf message if needed + pb_request = getattr(request, "_pb", request) + + transcoded_request = path_template.transcode(http_options, pb_request) + + body_json = None + if transcoded_request.get("body") is not None: + body_json = json_format.MessageToJson( + transcoded_request["body"], + use_integers_for_enums=rest_numeric_enums, + ) + + query_params_json = {} + if transcoded_request.get("query_params") is not None: + query_params_json = json.loads( + json_format.MessageToJson( + transcoded_request["query_params"], + use_integers_for_enums=rest_numeric_enums, + ) + ) + + # If required_fields_default_values is provided, we merge default values for missing + # required fields into the query parameters. + if required_fields_default_values: + for k, v in required_fields_default_values.items(): + if k not in query_params_json: + query_params_json[k] = v + + if rest_numeric_enums: + query_params_json["$alt"] = "json;enum-encoding=int" + + return transcoded_request, body_json, query_params_json diff --git a/google/cloud/aiplatform_v1beta1/services/dataset_service/client.py b/google/cloud/aiplatform_v1beta1/services/dataset_service/client.py index 6c8fa92922..fa0b56a813 100644 --- a/google/cloud/aiplatform_v1beta1/services/dataset_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/dataset_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -146,77 +153,13 @@ class DatasetServiceClient(metaclass=DatasetServiceClientMeta): resources. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -614,7 +557,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = DatasetServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -635,35 +578,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = DatasetServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -683,65 +603,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = DatasetServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = DatasetServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -871,13 +732,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - DatasetServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = DatasetServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = DatasetServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=DatasetServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -913,14 +776,14 @@ def __init__( self._transport = cast(DatasetServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or DatasetServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=DatasetServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=DatasetServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest.py index f2ccc6d221..041c556c29 100644 --- a/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3917,19 +3918,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseAssembleData._get_http_options() ) - request, metadata = self._interceptor.pre_assemble_data(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseAssembleData._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseAssembleData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseAssembleData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseAssembleData, + "_BaseAssembleData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4071,23 +4069,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseAssessData._get_http_options() ) - request, metadata = self._interceptor.pre_assess_data(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseAssessData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseAssessData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseAssessData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseAssessData, + "_BaseAssessData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4228,19 +4219,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDataset._get_http_options() ) - request, metadata = self._interceptor.pre_create_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDataset, + "_BaseCreateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4382,21 +4370,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_create_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion, + "_BaseCreateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4536,15 +4521,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_http_options() ) - request, metadata = self._interceptor.pre_delete_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDataset, + "_BaseDeleteDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4684,17 +4670,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_delete_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion, + "_BaseDeleteDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4833,17 +4820,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_http_options() ) - request, metadata = self._interceptor.pre_delete_saved_query( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery, + "_BaseDeleteSavedQuery__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4983,23 +4971,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseExportData._get_http_options() ) - request, metadata = self._interceptor.pre_export_data(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseExportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseExportData, + "_BaseExportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5138,17 +5119,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_http_options() ) - request, metadata = self._interceptor.pre_get_annotation_spec( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec, + "_BaseGetAnnotationSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5288,17 +5270,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDataset._get_http_options() ) - request, metadata = self._interceptor.pre_get_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseGetDataset._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDataset, + "_BaseGetDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5436,17 +5417,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_get_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDatasetVersion, + "_BaseGetDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5588,23 +5570,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseImportData._get_http_options() ) - request, metadata = self._interceptor.pre_import_data(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseImportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseImportData, + "_BaseImportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5743,17 +5718,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListAnnotations._get_http_options() ) - request, metadata = self._interceptor.pre_list_annotations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListAnnotations, + "_BaseListAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5895,15 +5871,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDataItems._get_http_options() ) - request, metadata = self._interceptor.pre_list_data_items(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDataItems, + "_BaseListDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6045,15 +6022,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasets._get_http_options() ) - request, metadata = self._interceptor.pre_list_datasets(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasets._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasets._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasets, + "_BaseListDatasets__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6196,17 +6174,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_http_options() ) - request, metadata = self._interceptor.pre_list_dataset_versions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasetVersions, + "_BaseListDatasetVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6348,17 +6327,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_http_options() ) - request, metadata = self._interceptor.pre_list_saved_queries( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListSavedQueries, + "_BaseListSavedQueries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6502,17 +6482,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_restore_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion, + "_BaseRestoreDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6650,17 +6631,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_http_options() ) - request, metadata = self._interceptor.pre_search_data_items( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSearchDataItems, + "_BaseSearchDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6803,19 +6785,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_http_options() ) - request, metadata = self._interceptor.pre_update_dataset(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDataset, + "_BaseUpdateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6956,21 +6935,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_http_options() ) - request, metadata = self._interceptor.pre_update_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion, + "_BaseUpdateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7310,15 +7286,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7451,15 +7428,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7593,19 +7571,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7740,19 +7715,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7887,21 +7859,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8032,17 +8001,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8147,17 +8117,18 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8265,15 +8236,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8406,15 +8378,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8547,15 +8520,16 @@ def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_asyncio.py index 591722e1fc..5696dfc072 100644 --- a/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1851,21 +1853,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseAssembleData._get_http_options() ) - request, metadata = await self._interceptor.pre_assemble_data( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseAssembleData._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseAssembleData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseAssembleData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseAssembleData, + "_BaseAssembleData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2017,25 +2016,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseAssessData._get_http_options() ) - request, metadata = await self._interceptor.pre_assess_data( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseAssessData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseAssessData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseAssessData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseAssessData, + "_BaseAssessData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2184,21 +2176,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_create_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDataset, + "_BaseCreateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2350,21 +2339,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_create_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion, + "_BaseCreateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2512,17 +2498,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDataset, + "_BaseDeleteDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2672,17 +2659,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion, + "_BaseDeleteDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2830,17 +2818,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_saved_query( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery, + "_BaseDeleteSavedQuery__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2990,25 +2979,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseExportData._get_http_options() ) - request, metadata = await self._interceptor.pre_export_data( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseExportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseExportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseExportData, + "_BaseExportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3156,17 +3138,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_http_options() ) - request, metadata = await self._interceptor.pre_get_annotation_spec( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec, + "_BaseGetAnnotationSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3314,19 +3297,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_get_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDataset._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseGetDataset._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDataset, + "_BaseGetDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3471,17 +3453,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_get_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetDatasetVersion, + "_BaseGetDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3631,25 +3614,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseImportData._get_http_options() ) - request, metadata = await self._interceptor.pre_import_data( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseImportData._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseDatasetServiceRestTransport._BaseImportData._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseImportData, + "_BaseImportData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3797,17 +3773,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListAnnotations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_annotations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListAnnotations, + "_BaseListAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3957,17 +3934,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDataItems._get_http_options() ) - request, metadata = await self._interceptor.pre_list_data_items( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDataItems, + "_BaseListDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4117,17 +4095,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasets._get_http_options() ) - request, metadata = await self._interceptor.pre_list_datasets( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasets._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasets._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasets, + "_BaseListDatasets__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4278,17 +4257,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_dataset_versions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListDatasetVersions, + "_BaseListDatasetVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4437,17 +4417,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_http_options() ) - request, metadata = await self._interceptor.pre_list_saved_queries( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListSavedQueries, + "_BaseListSavedQueries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4599,17 +4580,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_restore_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion, + "_BaseRestoreDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4758,17 +4740,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_http_options() ) - request, metadata = await self._interceptor.pre_search_data_items( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSearchDataItems, + "_BaseSearchDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4919,21 +4902,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_update_dataset( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDataset, + "_BaseUpdateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5082,21 +5062,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_update_dataset_version( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion, + "_BaseUpdateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7625,17 +7602,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7776,17 +7754,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7928,21 +7907,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8085,21 +8061,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8243,21 +8216,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8395,17 +8365,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8519,17 +8490,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8645,17 +8617,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8797,17 +8770,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8948,17 +8922,18 @@ async def __call__( http_options = ( _BaseDatasetServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDatasetServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDatasetServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_base.py index 0e4993e976..c66389ceb9 100644 --- a/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/dataset_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,52 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.AssembleDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseAssembleData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAssessData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -180,52 +132,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.AssessDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseAssessData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -242,52 +154,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.CreateDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseCreateDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -304,52 +176,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.CreateDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseCreateDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -364,43 +196,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.DeleteDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseDeleteDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -415,43 +216,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.DeleteDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseDeleteDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSavedQuery: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -462,43 +232,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.DeleteSavedQueryRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseDeleteSavedQuery._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -510,52 +249,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ExportDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseExportData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetAnnotationSpec: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -566,43 +265,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.GetAnnotationSpecRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseGetAnnotationSpec._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -617,43 +285,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.GetDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseGetDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -668,43 +305,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.GetDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseGetDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -716,52 +322,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ImportDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseImportData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListAnnotations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -772,43 +338,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListAnnotationsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListAnnotations._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDataItems: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -819,43 +354,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListDataItemsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListDataItems._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDatasets: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -870,43 +374,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListDatasetsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListDatasets._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDatasetVersions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -921,43 +394,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListDatasetVersionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListDatasetVersions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSavedQueries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -968,43 +410,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.ListSavedQueriesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseListSavedQueries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRestoreDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1019,43 +430,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.RestoreDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseRestoreDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchDataItems: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1066,29 +446,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.SearchDataItemsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseSearchDataItems._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1097,14 +454,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1121,38 +470,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.UpdateDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseUpdateDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateDatasetVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1161,14 +478,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1185,38 +494,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = dataset_service.UpdateDatasetVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDatasetServiceRestTransport._BaseUpdateDatasetVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1235,17 +512,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1264,17 +530,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1358,22 +613,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1464,22 +703,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1555,22 +778,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1981,17 +1188,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2446,17 +1642,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2919,17 +2104,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3388,17 +2562,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3853,16 +3016,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseDatasetServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/client.py b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/client.py index e0f2cb07a1..dbe556f715 100644 --- a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -149,77 +156,13 @@ class DeploymentResourcePoolServiceClient( ): """A service that manages the DeploymentResourcePool resource.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -473,9 +416,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - DeploymentResourcePoolServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -496,37 +437,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - DeploymentResourcePoolServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -546,67 +462,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -740,17 +595,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - DeploymentResourcePoolServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( DeploymentResourcePoolServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - DeploymentResourcePoolServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -788,14 +643,14 @@ def __init__( self._transport = cast(DeploymentResourcePoolServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or DeploymentResourcePoolServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest.py index 4c557de02d..932169f271 100644 --- a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3076,21 +3077,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_create_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool, + "_BaseCreateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3236,17 +3234,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_delete_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool, + "_BaseDeleteDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3392,17 +3391,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_get_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool, + "_BaseGetDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3550,17 +3550,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_http_options() ) - request, metadata = self._interceptor.pre_list_deployment_resource_pools( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools, + "_BaseListDeploymentResourcePools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3713,17 +3714,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_http_options() ) - request, metadata = self._interceptor.pre_query_deployed_models( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels, + "_BaseQueryDeployedModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3873,21 +3875,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_http_options() ) - request, metadata = self._interceptor.pre_update_deployment_resource_pool( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool, + "_BaseUpdateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4095,15 +4094,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4239,15 +4239,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4384,19 +4385,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4534,19 +4532,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4684,21 +4679,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4830,17 +4822,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4946,17 +4939,18 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5065,15 +5059,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5209,15 +5204,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5351,15 +5347,16 @@ def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_asyncio.py index 0215014eeb..975458ac03 100644 --- a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -933,23 +935,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool, + "_BaseCreateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1103,19 +1102,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool, + "_BaseDeleteDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1269,19 +1269,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool, + "_BaseGetDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1437,19 +1438,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_deployment_resource_pools( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools, + "_BaseListDeploymentResourcePools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1608,17 +1610,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_http_options() ) - request, metadata = await self._interceptor.pre_query_deployed_models( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels, + "_BaseQueryDeployedModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1774,23 +1777,20 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_deployment_resource_pool( request, metadata ) ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool, + "_BaseUpdateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4219,17 +4219,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4369,17 +4370,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4520,21 +4522,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4676,21 +4675,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4834,21 +4830,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4988,17 +4981,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5112,17 +5106,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5237,17 +5232,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5389,17 +5385,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5539,17 +5536,18 @@ async def __call__( http_options = ( _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseDeploymentResourcePoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_base.py index f0b8f213f7..310453c507 100644 --- a/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/deployment_resource_pool_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,56 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.CreateDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseCreateDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDeploymentResourcePool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -181,47 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.DeleteDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseDeleteDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDeploymentResourcePool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -232,47 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.GetDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseGetDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDeploymentResourcePools: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -283,47 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.ListDeploymentResourcePoolsRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseListDeploymentResourcePools._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryDeployedModels: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -334,31 +177,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = deployment_resource_pool_service.QueryDeployedModelsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseQueryDeployedModels._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateDeploymentResourcePool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -367,14 +185,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -386,42 +196,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - deployment_resource_pool_service.UpdateDeploymentResourcePoolRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseDeploymentResourcePoolServiceRestTransport._BaseUpdateDeploymentResourcePool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -440,17 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -469,17 +232,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -563,22 +315,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -669,22 +405,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -760,22 +480,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1186,17 +890,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1651,17 +1344,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2124,17 +1806,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2593,17 +2264,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3058,16 +2718,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseDeploymentResourcePoolServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/endpoint_service/client.py b/google/cloud/aiplatform_v1beta1/services/endpoint_service/client.py index 19f74c07f6..3b390541a3 100644 --- a/google/cloud/aiplatform_v1beta1/services/endpoint_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/endpoint_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class EndpointServiceClient(metaclass=EndpointServiceClientMeta): """A service for managing Vertex AI's Endpoints.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -503,7 +446,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = EndpointServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -524,35 +467,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = EndpointServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -572,65 +492,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = EndpointServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = EndpointServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -762,13 +623,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - EndpointServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = EndpointServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = EndpointServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=EndpointServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -804,14 +667,14 @@ def __init__( self._transport = cast(EndpointServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or EndpointServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=EndpointServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=EndpointServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest.py index ba90ff1238..dc70dca3f6 100644 --- a/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3337,19 +3338,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_create_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCreateEndpoint, + "_BaseCreateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3489,15 +3487,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_delete_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteEndpoint, + "_BaseDeleteEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3637,19 +3636,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeployModel._get_http_options() ) - request, metadata = self._interceptor.pre_deploy_model(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseDeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeployModel, + "_BaseDeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3790,17 +3786,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig._get_http_options() ) - request, metadata = self._interceptor.pre_fetch_publisher_model_config( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig, + "_BaseFetchPublisherModelConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3943,15 +3940,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_get_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetEndpoint, + "_BaseGetEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4091,15 +4089,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListEndpoints._get_http_options() ) - request, metadata = self._interceptor.pre_list_endpoints(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListEndpoints, + "_BaseListEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4244,21 +4243,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_http_options() ) - request, metadata = self._interceptor.pre_mutate_deployed_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseMutateDeployedModel, + "_BaseMutateDeployedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4401,21 +4397,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_http_options() ) - request, metadata = self._interceptor.pre_set_publisher_model_config( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig, + "_BaseSetPublisherModelConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4558,19 +4551,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUndeployModel._get_http_options() ) - request, metadata = self._interceptor.pre_undeploy_model(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUndeployModel, + "_BaseUndeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4711,19 +4701,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_update_endpoint(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpoint, + "_BaseUpdateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4868,21 +4855,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_http_options() ) - request, metadata = self._interceptor.pre_update_endpoint_long_running( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning, + "_BaseUpdateEndpointLongRunning__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5122,15 +5106,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5263,15 +5248,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5405,19 +5391,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5552,19 +5535,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5700,21 +5680,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5845,17 +5822,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5960,17 +5938,18 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6078,15 +6057,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6219,15 +6199,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6360,15 +6341,16 @@ def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_asyncio.py index 384579f0b1..85406e2506 100644 --- a/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1220,21 +1222,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_create_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCreateEndpoint, + "_BaseCreateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1385,17 +1384,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteEndpoint, + "_BaseDeleteEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1545,21 +1545,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeployModel._get_http_options() ) - request, metadata = await self._interceptor.pre_deploy_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseDeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeployModel, + "_BaseDeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1710,19 +1707,20 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig._get_http_options() ) - request, metadata = ( await self._interceptor.pre_fetch_publisher_model_config( request, metadata ) ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig, + "_BaseFetchPublisherModelConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1871,17 +1869,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_get_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetEndpoint, + "_BaseGetEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2030,17 +2029,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListEndpoints._get_http_options() ) - request, metadata = await self._interceptor.pre_list_endpoints( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListEndpoints, + "_BaseListEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2193,21 +2193,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_http_options() ) - request, metadata = await self._interceptor.pre_mutate_deployed_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseMutateDeployedModel, + "_BaseMutateDeployedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2358,21 +2355,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_http_options() ) - request, metadata = await self._interceptor.pre_set_publisher_model_config( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig, + "_BaseSetPublisherModelConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2524,21 +2518,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUndeployModel._get_http_options() ) - request, metadata = await self._interceptor.pre_undeploy_model( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUndeployModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUndeployModel, + "_BaseUndeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2690,21 +2681,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_update_endpoint( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpoint, + "_BaseUpdateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2857,23 +2845,20 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_endpoint_long_running( request, metadata ) ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning, + "_BaseUpdateEndpointLongRunning__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5324,17 +5309,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5476,17 +5462,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5629,21 +5616,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5787,21 +5771,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5945,21 +5926,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6097,17 +6075,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6221,17 +6200,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6348,17 +6328,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6500,17 +6481,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6652,17 +6634,18 @@ async def __call__( http_options = ( _BaseEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_base.py index da0b670545..0d88f18b7b 100644 --- a/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/endpoint_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,52 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.CreateEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseCreateEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -176,43 +128,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.DeleteEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseDeleteEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeployModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,52 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.DeployModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseDeployModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseFetchPublisherModelConfig: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -280,43 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.FetchPublisherModelConfigRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseFetchPublisherModelConfig._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -327,43 +177,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.GetEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseGetEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListEndpoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -374,43 +193,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.ListEndpointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseListEndpoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseMutateDeployedModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -422,52 +210,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.MutateDeployedModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseMutateDeployedModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSetPublisherModelConfig: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -479,52 +227,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.SetPublisherModelConfigRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseSetPublisherModelConfig._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUndeployModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -536,38 +244,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.UndeployModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseUndeployModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -576,14 +252,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -595,52 +263,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.UpdateEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseUpdateEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateEndpointLongRunning: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -652,38 +280,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = endpoint_service.UpdateEndpointLongRunningRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEndpointServiceRestTransport._BaseUpdateEndpointLongRunning._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -702,17 +298,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -731,17 +316,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -825,22 +399,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -931,22 +489,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1022,22 +564,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1448,17 +974,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1913,17 +1428,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2386,17 +1890,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2855,17 +2348,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3320,16 +2802,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseEndpointServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/evaluation_service/client.py b/google/cloud/aiplatform_v1beta1/services/evaluation_service/client.py index 6130f3e8b5..4d16f1a464 100644 --- a/google/cloud/aiplatform_v1beta1/services/evaluation_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/evaluation_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -130,77 +137,13 @@ def get_transport_class( class EvaluationServiceClient(metaclass=EvaluationServiceClientMeta): """Vertex AI Online Evaluation Service.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -388,7 +331,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = EvaluationServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -409,35 +352,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = EvaluationServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -457,65 +377,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = EvaluationServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = EvaluationServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -649,13 +510,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - EvaluationServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = EvaluationServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = EvaluationServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=EvaluationServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -691,14 +554,14 @@ def __init__( self._transport = cast(EvaluationServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or EvaluationServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=EvaluationServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest.py index a4550d5c0d..c1ab64e2ea 100644 --- a/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2893,21 +2894,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_http_options() ) - request, metadata = self._interceptor.pre_evaluate_dataset( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseEvaluateDataset, + "_BaseEvaluateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3048,21 +3046,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_http_options() ) - request, metadata = self._interceptor.pre_evaluate_instances( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseEvaluateInstances, + "_BaseEvaluateInstances__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3207,21 +3202,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_http_options() ) - request, metadata = self._interceptor.pre_generate_instance_rubrics( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics, + "_BaseGenerateInstanceRubrics__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3400,15 +3392,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3542,15 +3535,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3684,19 +3678,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3831,19 +3822,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3979,21 +3967,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4125,17 +4110,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4241,17 +4227,18 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4359,15 +4346,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4501,15 +4489,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4643,15 +4632,16 @@ def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_asyncio.py index a49b0b2c39..7d58d5b4fe 100644 --- a/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -735,21 +737,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_evaluate_dataset( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseEvaluateDataset, + "_BaseEvaluateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -898,21 +897,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_http_options() ) - request, metadata = await self._interceptor.pre_evaluate_instances( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseEvaluateInstances, + "_BaseEvaluateInstances__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1063,21 +1059,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_http_options() ) - request, metadata = await self._interceptor.pre_generate_instance_rubrics( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics, + "_BaseGenerateInstanceRubrics__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3482,17 +3475,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3634,17 +3628,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3787,21 +3782,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3945,21 +3937,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4103,21 +4092,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4255,17 +4241,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4377,17 +4364,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4502,17 +4490,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4654,17 +4643,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4806,17 +4796,18 @@ async def __call__( http_options = ( _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseEvaluationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseEvaluationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_base.py index f6b7da35b9..538cccdd21 100644 --- a/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/evaluation_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,52 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = evaluation_service.EvaluateDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEvaluationServiceRestTransport._BaseEvaluateDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseEvaluateInstances: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -185,52 +137,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = evaluation_service.EvaluateInstancesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEvaluationServiceRestTransport._BaseEvaluateInstances._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGenerateInstanceRubrics: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -247,38 +159,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = evaluation_service.GenerateInstanceRubricsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseEvaluationServiceRestTransport._BaseGenerateInstanceRubrics._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -297,17 +177,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -326,17 +195,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -420,22 +278,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -526,22 +368,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -617,22 +443,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1043,17 +853,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1508,17 +1307,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1981,17 +1769,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2450,17 +2227,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2915,16 +2681,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseEvaluationServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/example_store_service/client.py b/google/cloud/aiplatform_v1beta1/services/example_store_service/client.py index 830bd31c94..2a617e5d9e 100644 --- a/google/cloud/aiplatform_v1beta1/services/example_store_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/example_store_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -135,77 +142,13 @@ def get_transport_class( class ExampleStoreServiceClient(metaclass=ExampleStoreServiceClientMeta): """A service for managing and retrieving few-shot examples.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -393,7 +336,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ExampleStoreServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -414,35 +357,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ExampleStoreServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -462,65 +382,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ExampleStoreServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ExampleStoreServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ExampleStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ExampleStoreServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -654,13 +515,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ExampleStoreServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ExampleStoreServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ExampleStoreServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ExampleStoreServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -696,14 +559,14 @@ def __init__( self._transport = cast(ExampleStoreServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ExampleStoreServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ExampleStoreServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ExampleStoreServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ExampleStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest.py index 9e965de0f8..b48a304889 100644 --- a/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3245,21 +3246,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_http_options() ) - request, metadata = self._interceptor.pre_create_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore, + "_BaseCreateExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3402,17 +3400,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore._get_http_options() ) - request, metadata = self._interceptor.pre_delete_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore, + "_BaseDeleteExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3554,19 +3553,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_http_options() ) - request, metadata = self._interceptor.pre_fetch_examples(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseFetchExamples, + "_BaseFetchExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3710,17 +3706,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetExampleStore._get_http_options() ) - request, metadata = self._interceptor.pre_get_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetExampleStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetExampleStore, + "_BaseGetExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3861,17 +3858,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseListExampleStores._get_http_options() ) - request, metadata = self._interceptor.pre_list_example_stores( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseListExampleStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseListExampleStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseListExampleStores, + "_BaseListExampleStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4019,19 +4017,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_http_options() ) - request, metadata = self._interceptor.pre_remove_examples(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseRemoveExamples, + "_BaseRemoveExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4176,19 +4171,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_http_options() ) - request, metadata = self._interceptor.pre_search_examples(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseSearchExamples, + "_BaseSearchExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4334,21 +4326,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_http_options() ) - request, metadata = self._interceptor.pre_update_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore, + "_BaseUpdateExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4491,19 +4480,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_http_options() ) - request, metadata = self._interceptor.pre_upsert_examples(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseUpsertExamples, + "_BaseUpsertExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4742,15 +4728,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4884,15 +4871,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5027,19 +5015,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5175,19 +5160,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5323,21 +5305,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5471,17 +5450,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5587,17 +5567,18 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5706,15 +5687,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5848,15 +5830,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5990,15 +5973,16 @@ def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_asyncio.py index c4a6297b4e..90d710a6f9 100644 --- a/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1117,21 +1119,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_http_options() ) - request, metadata = await self._interceptor.pre_create_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore, + "_BaseCreateExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1280,17 +1279,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore, + "_BaseDeleteExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1438,21 +1438,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_http_options() ) - request, metadata = await self._interceptor.pre_fetch_examples( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseFetchExamples, + "_BaseFetchExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1602,17 +1599,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetExampleStore._get_http_options() ) - request, metadata = await self._interceptor.pre_get_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetExampleStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetExampleStore, + "_BaseGetExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1759,17 +1757,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseListExampleStores._get_http_options() ) - request, metadata = await self._interceptor.pre_list_example_stores( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseListExampleStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseListExampleStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseListExampleStores, + "_BaseListExampleStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1921,21 +1920,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_http_options() ) - request, metadata = await self._interceptor.pre_remove_examples( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseRemoveExamples, + "_BaseRemoveExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2086,21 +2082,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_http_options() ) - request, metadata = await self._interceptor.pre_search_examples( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseSearchExamples, + "_BaseSearchExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2252,21 +2245,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_http_options() ) - request, metadata = await self._interceptor.pre_update_example_store( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore, + "_BaseUpdateExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2415,21 +2405,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_http_options() ) - request, metadata = await self._interceptor.pre_upsert_examples( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseUpsertExamples, + "_BaseUpsertExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4881,17 +4868,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5033,17 +5021,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5184,21 +5173,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5342,21 +5328,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5500,21 +5483,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5652,17 +5632,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5774,17 +5755,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5899,17 +5881,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6051,17 +6034,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6201,17 +6185,18 @@ async def __call__( http_options = ( _BaseExampleStoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseExampleStoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExampleStoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExampleStoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_base.py index a2425e4432..e373098aed 100644 --- a/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/example_store_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,52 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.CreateExampleStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseCreateExampleStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteExampleStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -179,43 +131,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.DeleteExampleStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseDeleteExampleStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseFetchExamples: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -227,52 +148,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.FetchExamplesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseFetchExamples._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetExampleStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -283,43 +164,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.GetExampleStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseGetExampleStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListExampleStores: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -330,43 +180,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.ListExampleStoresRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseListExampleStores._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRemoveExamples: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -378,52 +197,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.RemoveExamplesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseRemoveExamples._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchExamples: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -435,52 +214,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.SearchExamplesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseSearchExamples._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateExampleStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -492,52 +231,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.UpdateExampleStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseUpdateExampleStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpsertExamples: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -549,38 +248,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = example_store_service.UpsertExamplesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExampleStoreServiceRestTransport._BaseUpsertExamples._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -599,17 +266,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -628,17 +284,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -722,22 +367,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -828,22 +457,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -919,22 +532,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1345,17 +942,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1810,17 +1396,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2283,17 +1858,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2752,17 +2316,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3217,16 +2770,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseExampleStoreServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/client.py b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/client.py index bcf12229db..e36bc4bd3b 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -130,77 +137,13 @@ def get_transport_class( class ExtensionExecutionServiceClient(metaclass=ExtensionExecutionServiceClientMeta): """A service for Extension execution.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -410,7 +353,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ExtensionExecutionServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -431,35 +374,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ExtensionExecutionServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -479,67 +399,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ExtensionExecutionServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ExtensionExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - ExtensionExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ExtensionExecutionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -673,15 +532,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ExtensionExecutionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( ExtensionExecutionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ExtensionExecutionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ExtensionExecutionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -717,14 +578,14 @@ def __init__( self._transport = cast(ExtensionExecutionServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ExtensionExecutionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ExtensionExecutionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ExtensionExecutionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ExtensionExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest.py index 1c7365ee08..f08b70db01 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -609,21 +610,18 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_http_options() ) - request, metadata = self._interceptor.pre_execute_extension( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension, + "_BaseExecuteExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -772,19 +770,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_http_options() ) - request, metadata = self._interceptor.pre_query_extension(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension, + "_BaseQueryExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -954,15 +949,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1098,15 +1094,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1243,19 +1240,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1393,19 +1387,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1543,21 +1534,18 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1689,17 +1677,18 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1807,17 +1796,18 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1928,15 +1918,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2072,15 +2063,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2216,15 +2208,16 @@ def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_asyncio.py index 104326a1fd..d665384327 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -671,21 +673,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_http_options() ) - request, metadata = await self._interceptor.pre_execute_extension( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension, + "_BaseExecuteExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -838,21 +837,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_http_options() ) - request, metadata = await self._interceptor.pre_query_extension( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension, + "_BaseQueryExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1022,17 +1018,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1172,17 +1169,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1323,21 +1321,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1479,21 +1474,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1637,21 +1629,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1789,17 +1778,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1911,17 +1901,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2036,17 +2027,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2186,17 +2178,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2336,17 +2329,18 @@ async def __call__( http_options = ( _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionExecutionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_base.py index 7b3aa3bb6c..f871065a9c 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_execution_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,52 +110,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = extension_execution_service.ExecuteExtensionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExtensionExecutionServiceRestTransport._BaseExecuteExtension._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryExtension: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,38 +127,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = extension_execution_service.QueryExtensionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExtensionExecutionServiceRestTransport._BaseQueryExtension._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -225,17 +145,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -254,17 +163,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -348,22 +246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -454,22 +336,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -545,22 +411,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -971,17 +821,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1436,17 +1275,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1909,17 +1737,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2378,17 +2195,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2843,16 +2649,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseExtensionExecutionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/client.py b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/client.py index 0efeae99e0..025327ec82 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -137,77 +144,13 @@ def get_transport_class( class ExtensionRegistryServiceClient(metaclass=ExtensionRegistryServiceClientMeta): """A service for managing Vertex AI's Extension registry.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -441,7 +384,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ExtensionRegistryServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -462,35 +405,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ExtensionRegistryServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -510,67 +430,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ExtensionRegistryServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ExtensionRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - ExtensionRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ExtensionRegistryServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -704,15 +563,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ExtensionRegistryServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( ExtensionRegistryServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ExtensionRegistryServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ExtensionRegistryServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -748,14 +609,14 @@ def __init__( self._transport = cast(ExtensionRegistryServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ExtensionRegistryServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ExtensionRegistryServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ExtensionRegistryServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ExtensionRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest.py index 0330fd9422..ab46fedbdb 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3003,17 +3004,18 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension._get_http_options() ) - request, metadata = self._interceptor.pre_delete_extension( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension, + "_BaseDeleteExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3155,15 +3157,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetExtension._get_http_options() ) - request, metadata = self._interceptor.pre_get_extension(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetExtension._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetExtension, + "_BaseGetExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3308,21 +3311,18 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_http_options() ) - request, metadata = self._interceptor.pre_import_extension( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseImportExtension, + "_BaseImportExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3464,15 +3464,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseListExtensions._get_http_options() ) - request, metadata = self._interceptor.pre_list_extensions(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseListExtensions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseListExtensions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseListExtensions, + "_BaseListExtensions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3621,21 +3622,18 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_http_options() ) - request, metadata = self._interceptor.pre_update_extension( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension, + "_BaseUpdateExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3830,15 +3828,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3972,15 +3971,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4117,19 +4117,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4267,19 +4264,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4417,21 +4411,18 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4565,17 +4556,18 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4683,17 +4675,18 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4804,15 +4797,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4948,15 +4942,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5092,15 +5087,16 @@ def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_asyncio.py index 355efc2be2..181c0520bb 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -861,17 +863,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_extension( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension, + "_BaseDeleteExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1019,17 +1022,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetExtension._get_http_options() ) - request, metadata = await self._interceptor.pre_get_extension( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetExtension._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetExtension, + "_BaseGetExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1178,21 +1182,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_http_options() ) - request, metadata = await self._interceptor.pre_import_extension( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseImportExtension, + "_BaseImportExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1340,17 +1341,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseListExtensions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_extensions( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseListExtensions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseListExtensions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseListExtensions, + "_BaseListExtensions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1503,21 +1505,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_http_options() ) - request, metadata = await self._interceptor.pre_update_extension( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension, + "_BaseUpdateExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3931,17 +3930,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4081,17 +4081,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4232,21 +4233,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4388,21 +4386,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4544,21 +4539,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4696,17 +4688,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4818,17 +4811,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4943,17 +4937,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5093,17 +5088,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5243,17 +5239,18 @@ async def __call__( http_options = ( _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseExtensionRegistryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_base.py index 58b1d25cfd..fe29e0cf98 100644 --- a/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/extension_registry_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,43 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = extension_registry_service.DeleteExtensionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExtensionRegistryServiceRestTransport._BaseDeleteExtension._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetExtension: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -166,43 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = extension_registry_service.GetExtensionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExtensionRegistryServiceRestTransport._BaseGetExtension._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportExtension: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -214,52 +144,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = extension_registry_service.ImportExtensionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExtensionRegistryServiceRestTransport._BaseImportExtension._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListExtensions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -270,29 +160,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = extension_registry_service.ListExtensionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExtensionRegistryServiceRestTransport._BaseListExtensions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateExtension: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -301,14 +168,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -320,38 +179,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = extension_registry_service.UpdateExtensionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseExtensionRegistryServiceRestTransport._BaseUpdateExtension._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -370,17 +197,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -399,17 +215,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -493,22 +298,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -599,22 +388,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -690,22 +463,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1116,17 +873,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1581,17 +1327,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2054,17 +1789,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2523,17 +2247,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2988,16 +2701,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseExtensionRegistryServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/client.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/client.py index 0401970f07..1c77ce8fbb 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -159,77 +166,13 @@ class FeatureOnlineStoreAdminServiceClient( FeatureOnlineStore. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -465,9 +408,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -488,37 +429,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -538,67 +454,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -732,17 +587,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeatureOnlineStoreAdminServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( FeatureOnlineStoreAdminServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - FeatureOnlineStoreAdminServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -780,14 +635,14 @@ def __init__( self._transport = cast(FeatureOnlineStoreAdminServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeatureOnlineStoreAdminServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest.py index 7f12b91e4d..217e3038a6 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3485,21 +3486,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore, + "_BaseCreateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3641,21 +3639,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView, + "_BaseCreateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3799,17 +3794,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore, + "_BaseDeleteFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3949,17 +3945,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView, + "_BaseDeleteFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4103,17 +4100,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore, + "_BaseGetFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4257,17 +4255,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView, + "_BaseGetFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4412,17 +4411,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_view_sync( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync, + "_BaseGetFeatureViewSync__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4568,17 +4568,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_online_stores( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores, + "_BaseListFeatureOnlineStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4725,17 +4726,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_views( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews, + "_BaseListFeatureViews__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4882,17 +4884,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_view_syncs( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs, + "_BaseListFeatureViewSyncs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5038,21 +5041,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_sync_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView, + "_BaseSyncFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5203,21 +5203,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore, + "_BaseUpdateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5359,21 +5356,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView, + "_BaseUpdateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5656,15 +5650,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5800,15 +5795,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5943,19 +5939,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6093,19 +6086,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6245,21 +6235,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6391,17 +6378,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6507,17 +6495,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6626,15 +6615,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6770,15 +6760,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6912,15 +6903,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_asyncio.py index 0391b6afdd..535f48f845 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1379,21 +1381,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore, + "_BaseCreateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1547,21 +1546,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView, + "_BaseCreateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1713,17 +1709,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore, + "_BaseDeleteFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1875,17 +1872,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView, + "_BaseDeleteFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2037,17 +2035,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore, + "_BaseGetFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2201,17 +2200,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView, + "_BaseGetFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2362,17 +2362,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_view_sync( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync, + "_BaseGetFeatureViewSync__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2524,17 +2525,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_online_stores( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores, + "_BaseListFeatureOnlineStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2691,17 +2693,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_views( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews, + "_BaseListFeatureViews__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2854,17 +2857,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_view_syncs( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs, + "_BaseListFeatureViewSyncs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3020,21 +3024,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_sync_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView, + "_BaseSyncFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3191,21 +3192,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature_online_store( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore, + "_BaseUpdateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3359,21 +3357,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature_view( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView, + "_BaseUpdateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5863,17 +5858,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6015,17 +6011,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6166,21 +6163,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6322,21 +6316,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6480,21 +6471,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6634,17 +6622,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6758,17 +6747,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6883,17 +6873,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7035,17 +7026,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7187,17 +7179,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_base.py index 333ebc969d..3e511bcd36 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_admin_service/transports/rest_base.py @@ -106,14 +106,6 @@ def __hash__(self): # pragma: NO COVER "featureOnlineStoreId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -125,42 +117,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.CreateFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -169,14 +125,6 @@ def __hash__(self): # pragma: NO COVER "featureViewId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -188,54 +136,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.CreateFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseCreateFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureOnlineStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -246,47 +152,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.DeleteFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -297,45 +168,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.DeleteFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseDeleteFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureOnlineStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -346,47 +184,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.GetFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -397,45 +200,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.GetFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureViewSync: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -446,45 +216,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.GetFeatureViewSyncRequest.pb(request) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseGetFeatureViewSync._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureOnlineStores: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -495,47 +232,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.ListFeatureOnlineStoresRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureOnlineStores._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureViews: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -546,45 +248,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.ListFeatureViewsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViews._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureViewSyncs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -595,47 +264,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.ListFeatureViewSyncsRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseListFeatureViewSyncs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSyncFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -647,54 +281,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.SyncFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseSyncFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeatureOnlineStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -706,56 +298,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_admin_service.UpdateFeatureOnlineStoreRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureOnlineStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeatureView: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -767,40 +315,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_admin_service.UpdateFeatureViewRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreAdminServiceRestTransport._BaseUpdateFeatureView._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -819,17 +333,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -848,17 +351,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -942,22 +434,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1048,22 +524,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1139,22 +599,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1565,17 +1009,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2030,17 +1463,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2503,17 +1925,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2972,17 +2383,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3437,16 +2837,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeatureOnlineStoreAdminServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/client.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/client.py index 386292d969..21104f1bb0 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/client.py @@ -41,6 +41,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -134,77 +141,13 @@ def get_transport_class( class FeatureOnlineStoreServiceClient(metaclass=FeatureOnlineStoreServiceClientMeta): """A service for fetching feature values from the online store.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -394,7 +337,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = FeatureOnlineStoreServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -415,35 +358,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = FeatureOnlineStoreServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -463,67 +383,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -657,15 +516,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeatureOnlineStoreServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( FeatureOnlineStoreServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = FeatureOnlineStoreServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -701,14 +562,14 @@ def __init__( self._transport = cast(FeatureOnlineStoreServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeatureOnlineStoreServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest.py index 9679dcd987..44d34c57df 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -690,21 +691,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_fetch_feature_values( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues, + "_BaseFetchFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -854,21 +852,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_http_options() ) - request, metadata = self._interceptor.pre_generate_fetch_access_token( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken, + "_BaseGenerateFetchAccessToken__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1015,21 +1010,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_http_options() ) - request, metadata = self._interceptor.pre_search_nearest_entities( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities, + "_BaseSearchNearestEntities__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1251,15 +1243,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1395,15 +1388,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1540,19 +1534,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1690,19 +1681,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1840,21 +1828,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1986,17 +1971,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2104,17 +2090,18 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2225,15 +2212,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2369,15 +2357,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2513,15 +2502,16 @@ def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_asyncio.py index 8cf2eb9d5b..88ad97565b 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -771,21 +773,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_fetch_feature_values( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues, + "_BaseFetchFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -941,21 +940,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_http_options() ) - request, metadata = await self._interceptor.pre_generate_fetch_access_token( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGenerateFetchAccessToken, + "_BaseGenerateFetchAccessToken__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1112,21 +1108,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_http_options() ) - request, metadata = await self._interceptor.pre_search_nearest_entities( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities, + "_BaseSearchNearestEntities__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1346,17 +1339,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1496,17 +1490,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1647,21 +1642,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1803,21 +1795,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1961,21 +1950,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2113,17 +2099,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2235,17 +2222,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2360,17 +2348,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2510,17 +2499,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2660,17 +2650,18 @@ async def __call__( http_options = ( _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureOnlineStoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_base.py index 5d0e04aef6..44d70ca208 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_online_store_service/transports/rest_base.py @@ -103,14 +103,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -122,40 +114,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_service.FetchFeatureValuesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreServiceRestTransport._BaseFetchFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGenerateFetchAccessToken: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -171,49 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - feature_online_store_service.GenerateFetchAccessTokenRequest.pb(request) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchNearestEntities: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -225,40 +146,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_online_store_service.SearchNearestEntitiesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureOnlineStoreServiceRestTransport._BaseSearchNearestEntities._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamingFetchFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -281,17 +168,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -310,17 +186,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -404,22 +269,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -510,22 +359,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -601,22 +434,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1027,17 +844,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1492,17 +1298,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1965,17 +1760,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2434,17 +2218,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2899,16 +2672,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeatureOnlineStoreServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/client.py b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/client.py index c971622c4b..5397704fa5 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -150,77 +157,13 @@ class FeatureRegistryServiceClient(metaclass=FeatureRegistryServiceClientMeta): FeatureRegistry. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -484,7 +427,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = FeatureRegistryServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -505,35 +448,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = FeatureRegistryServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -553,67 +473,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeatureRegistryServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeatureRegistryServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -747,13 +606,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeatureRegistryServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = FeatureRegistryServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = FeatureRegistryServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeatureRegistryServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -789,14 +650,14 @@ def __init__( self._transport = cast(FeatureRegistryServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeatureRegistryServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeatureRegistryServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest.py index e2c7d8ad1a..761fb30ee4 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3824,21 +3825,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3984,19 +3982,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4138,21 +4133,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup, + "_BaseCreateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4296,21 +4288,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor, + "_BaseCreateFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4452,21 +4441,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature_monitor_job( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob, + "_BaseCreateFeatureMonitorJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4613,15 +4599,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4761,17 +4748,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup, + "_BaseDeleteFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4913,17 +4901,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor, + "_BaseDeleteFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5067,15 +5056,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5214,17 +5204,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup, + "_BaseGetFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5365,17 +5356,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor, + "_BaseGetFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5516,17 +5508,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature_monitor_job( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob, + "_BaseGetFeatureMonitorJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5671,17 +5664,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_groups( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups, + "_BaseListFeatureGroups__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5828,17 +5822,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_monitor_jobs( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs, + "_BaseListFeatureMonitorJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5983,17 +5978,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors._get_http_options() ) - request, metadata = self._interceptor.pre_list_feature_monitors( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors, + "_BaseListFeatureMonitors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6144,15 +6140,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_list_features(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6299,19 +6296,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6453,21 +6447,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup, + "_BaseUpdateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6611,21 +6602,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor, + "_BaseUpdateFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6962,15 +6950,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7104,15 +7093,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7247,19 +7237,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7395,19 +7382,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7543,21 +7527,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7691,17 +7672,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7809,17 +7791,18 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7930,15 +7913,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8072,15 +8056,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8216,15 +8201,16 @@ def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_asyncio.py index 4ff454ddb1..462c6c3b5b 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1748,21 +1750,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1914,21 +1913,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2078,21 +2074,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup, + "_BaseCreateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2242,21 +2235,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor, + "_BaseCreateFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2406,21 +2396,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature_monitor_job( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob, + "_BaseCreateFeatureMonitorJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2575,17 +2562,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2733,17 +2721,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup, + "_BaseDeleteFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2891,17 +2880,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor, + "_BaseDeleteFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3051,17 +3041,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3206,17 +3197,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup, + "_BaseGetFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3361,17 +3353,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor, + "_BaseGetFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3516,17 +3509,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature_monitor_job( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob, + "_BaseGetFeatureMonitorJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3677,17 +3671,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_groups( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups, + "_BaseListFeatureGroups__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3840,17 +3835,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_monitor_jobs( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs, + "_BaseListFeatureMonitorJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4003,17 +3999,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors._get_http_options() ) - request, metadata = await self._interceptor.pre_list_feature_monitors( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors, + "_BaseListFeatureMonitors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4168,17 +4165,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_list_features( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4331,21 +4329,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4495,21 +4490,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature_group( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup, + "_BaseUpdateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4659,21 +4651,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature_monitor( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor, + "_BaseUpdateFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7203,17 +7192,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7353,17 +7343,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7504,21 +7495,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7660,21 +7648,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7816,21 +7801,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7968,17 +7950,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8090,17 +8073,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8215,17 +8199,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8365,17 +8350,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8515,17 +8501,18 @@ async def __call__( http_options = ( _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeatureRegistryServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_base.py index 0064f5a023..ffd4bcaed8 100644 --- a/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/feature_registry_service/transports/rest_base.py @@ -107,14 +107,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -126,38 +118,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.BatchCreateFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseBatchCreateFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -166,14 +126,6 @@ def __hash__(self): # pragma: NO COVER "featureId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -185,38 +137,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseCreateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -225,14 +145,6 @@ def __hash__(self): # pragma: NO COVER "featureGroupId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -244,38 +156,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.CreateFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeatureMonitor: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -284,14 +164,6 @@ def __hash__(self): # pragma: NO COVER "featureMonitorId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -303,54 +175,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.CreateFeatureMonitorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeatureMonitorJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -362,54 +192,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.CreateFeatureMonitorJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseCreateFeatureMonitorJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -420,43 +208,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -467,43 +224,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.DeleteFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureMonitor: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -514,45 +240,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.DeleteFeatureMonitorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseDeleteFeatureMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -563,43 +256,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseGetFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -610,43 +272,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.GetFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureMonitor: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -657,43 +288,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.GetFeatureMonitorRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeatureMonitorJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -704,45 +304,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.GetFeatureMonitorJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseGetFeatureMonitorJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureGroups: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -753,43 +320,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.ListFeatureGroupsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseListFeatureGroups._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureMonitorJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -800,45 +336,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.ListFeatureMonitorJobsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitorJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatureMonitors: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -849,43 +352,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.ListFeatureMonitorsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseListFeatureMonitors._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatures: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -896,43 +368,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseListFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -944,52 +385,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeatureGroup: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1001,52 +402,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.UpdateFeatureGroupRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureGroup._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeatureMonitor: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1058,40 +419,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = feature_registry_service.UpdateFeatureMonitorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeatureRegistryServiceRestTransport._BaseUpdateFeatureMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1110,17 +437,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1139,17 +455,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1233,22 +538,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1339,22 +628,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1430,22 +703,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1856,17 +1113,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2321,17 +1567,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2794,17 +2029,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3263,17 +2487,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3728,16 +2941,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeatureRegistryServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/client.py b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/client.py index 473ffce3ac..012f5df8d7 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/client.py @@ -40,6 +40,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -141,77 +148,13 @@ class FeaturestoreOnlineServingServiceClient( ): """A service for serving online feature values.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -401,9 +344,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - FeaturestoreOnlineServingServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -424,37 +365,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - FeaturestoreOnlineServingServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -474,65 +390,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -666,17 +523,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeaturestoreOnlineServingServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( FeaturestoreOnlineServingServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - FeaturestoreOnlineServingServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -714,14 +571,14 @@ def __init__( self._transport = cast(FeaturestoreOnlineServingServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeaturestoreOnlineServingServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest.py index b0383b768b..a9a906227e 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -672,21 +673,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues, + "_BaseReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -837,21 +835,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_streaming_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues, + "_BaseStreamingReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -992,21 +987,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_write_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues, + "_BaseWriteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1185,15 +1177,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1327,15 +1320,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1470,19 +1464,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1618,19 +1609,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1768,21 +1756,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1914,17 +1899,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2030,17 +2016,18 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2149,15 +2136,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2291,15 +2279,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2433,15 +2422,16 @@ def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_asyncio.py index ae92bcca2e..3148aee83f 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -740,21 +742,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues, + "_BaseReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -910,23 +909,20 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_http_options() ) - request, metadata = ( await self._interceptor.pre_streaming_read_feature_values( request, metadata ) ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues, + "_BaseStreamingReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1073,21 +1069,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_write_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues, + "_BaseWriteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1268,17 +1261,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1420,17 +1414,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1573,21 +1568,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1731,21 +1723,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1889,21 +1878,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2043,17 +2029,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2167,17 +2154,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2294,17 +2282,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2446,17 +2435,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2598,17 +2588,18 @@ async def __call__( http_options = ( _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_base.py index c2686b4fe3..7e3b58adaa 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_online_serving_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,54 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_online_service.ReadFeatureValuesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreOnlineServingServiceRestTransport._BaseReadFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamingReadFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -179,56 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - featurestore_online_service.StreamingReadFeatureValuesRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreOnlineServingServiceRestTransport._BaseStreamingReadFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseWriteFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -240,40 +146,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_online_service.WriteFeatureValuesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreOnlineServingServiceRestTransport._BaseWriteFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -292,17 +164,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -321,17 +182,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -415,22 +265,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -521,22 +355,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -612,22 +430,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1038,17 +840,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1503,17 +1294,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1976,17 +1756,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2445,17 +2214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2910,16 +2668,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeaturestoreOnlineServingServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_service/client.py b/google/cloud/aiplatform_v1beta1/services/featurestore_service/client.py index cc9d897350..e70e09cb7e 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -146,77 +153,13 @@ class FeaturestoreServiceClient(metaclass=FeaturestoreServiceClientMeta): Featurestore. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -454,7 +397,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = FeaturestoreServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -475,35 +418,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = FeaturestoreServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -523,65 +443,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = FeaturestoreServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = FeaturestoreServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -715,13 +576,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - FeaturestoreServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = FeaturestoreServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = FeaturestoreServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=FeaturestoreServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -757,14 +620,14 @@ def __init__( self._transport = cast(FeaturestoreServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or FeaturestoreServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=FeaturestoreServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest.py index 16b418fe8e..ae2e3b460d 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3928,21 +3929,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4086,21 +4084,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_batch_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues, + "_BaseBatchReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4244,21 +4239,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_create_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType, + "_BaseCreateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4402,19 +4394,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_create_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4556,21 +4545,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_create_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore, + "_BaseCreateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4713,17 +4699,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_delete_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType, + "_BaseDeleteEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4865,15 +4852,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5013,17 +5001,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_delete_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore, + "_BaseDeleteFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5166,21 +5155,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_delete_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues, + "_BaseDeleteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5324,21 +5310,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_export_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues, + "_BaseExportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5484,15 +5467,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_get_entity_type(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetEntityType, + "_BaseGetEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5636,15 +5620,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = self._interceptor.pre_get_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5788,17 +5773,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_get_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore, + "_BaseGetFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5941,21 +5927,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_http_options() ) - request, metadata = self._interceptor.pre_import_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues, + "_BaseImportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6097,17 +6080,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_http_options() ) - request, metadata = self._interceptor.pre_list_entity_types( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes, + "_BaseListEntityTypes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6254,15 +6238,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_list_features(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6405,17 +6390,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_http_options() ) - request, metadata = self._interceptor.pre_list_featurestores( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores, + "_BaseListFeaturestores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6560,15 +6546,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_http_options() ) - request, metadata = self._interceptor.pre_search_features(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures, + "_BaseSearchFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6716,21 +6703,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_http_options() ) - request, metadata = self._interceptor.pre_update_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType, + "_BaseUpdateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6876,19 +6860,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = self._interceptor.pre_update_feature(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7032,21 +7013,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_http_options() ) - request, metadata = self._interceptor.pre_update_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore, + "_BaseUpdateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7396,15 +7374,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7538,15 +7517,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7681,19 +7661,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7829,19 +7806,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7977,21 +7951,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8125,17 +8096,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8241,17 +8213,18 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8360,15 +8333,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8502,15 +8476,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8644,15 +8619,16 @@ def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_asyncio.py index 0140af3d96..9a0e5f1dba 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1862,21 +1864,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_create_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2026,21 +2025,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_read_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues, + "_BaseBatchReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2192,21 +2188,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_create_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType, + "_BaseCreateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2358,21 +2351,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_create_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2522,21 +2512,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_create_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore, + "_BaseCreateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2685,17 +2672,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType, + "_BaseDeleteEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2845,17 +2833,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3003,17 +2992,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore, + "_BaseDeleteFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3162,21 +3152,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues, + "_BaseDeleteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3326,21 +3313,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_export_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues, + "_BaseExportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3492,17 +3476,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_get_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetEntityType, + "_BaseGetEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3652,17 +3637,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_get_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3814,17 +3800,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_get_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore, + "_BaseGetFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3973,21 +3960,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_http_options() ) - request, metadata = await self._interceptor.pre_import_feature_values( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues, + "_BaseImportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4135,17 +4119,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_http_options() ) - request, metadata = await self._interceptor.pre_list_entity_types( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes, + "_BaseListEntityTypes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4298,17 +4283,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_list_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4459,17 +4445,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_http_options() ) - request, metadata = await self._interceptor.pre_list_featurestores( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores, + "_BaseListFeaturestores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4618,17 +4605,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_http_options() ) - request, metadata = await self._interceptor.pre_search_features( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures, + "_BaseSearchFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4782,21 +4770,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_http_options() ) - request, metadata = await self._interceptor.pre_update_entity_type( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType, + "_BaseUpdateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4948,21 +4933,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_http_options() ) - request, metadata = await self._interceptor.pre_update_feature( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5112,21 +5094,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_http_options() ) - request, metadata = await self._interceptor.pre_update_featurestore( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore, + "_BaseUpdateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7665,17 +7644,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7817,17 +7797,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7968,21 +7949,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8126,21 +8104,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8284,21 +8259,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8436,17 +8408,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8558,17 +8531,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8683,17 +8657,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8835,17 +8810,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8985,17 +8961,18 @@ async def __call__( http_options = ( _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseFeaturestoreServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseFeaturestoreServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_base.py index 1b30dafaa6..c09c76bac5 100644 --- a/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/featurestore_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,52 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.BatchCreateFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseBatchCreateFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchReadFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -180,38 +132,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.BatchReadFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseBatchReadFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -220,14 +140,6 @@ def __hash__(self): # pragma: NO COVER "entityTypeId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -239,38 +151,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseCreateEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -279,14 +159,6 @@ def __hash__(self): # pragma: NO COVER "featureId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -298,38 +170,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseCreateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -338,14 +178,6 @@ def __hash__(self): # pragma: NO COVER "featurestoreId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -357,52 +189,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.CreateFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseCreateFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -413,43 +205,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -460,43 +221,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -507,43 +237,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -555,52 +254,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.DeleteFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseDeleteFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -612,52 +271,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ExportFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseExportFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -668,43 +287,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseGetEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -715,43 +303,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseGetFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -762,43 +319,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.GetFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseGetFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportFeatureValues: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -810,52 +336,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ImportFeatureValuesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseImportFeatureValues._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListEntityTypes: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -866,43 +352,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListEntityTypesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseListEntityTypes._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeatures: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -913,43 +368,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseListFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListFeaturestores: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -960,43 +384,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.ListFeaturestoresRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseListFeaturestores._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchFeatures: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1007,43 +400,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.SearchFeaturesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseSearchFeatures._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateEntityType: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1055,52 +417,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateEntityTypeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseUpdateEntityType._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeature: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1112,52 +434,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateFeatureRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseUpdateFeature._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateFeaturestore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1169,38 +451,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = featurestore_service.UpdateFeaturestoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseFeaturestoreServiceRestTransport._BaseUpdateFeaturestore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1219,17 +469,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1248,17 +487,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1342,22 +570,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1448,22 +660,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1539,22 +735,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1965,17 +1145,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2430,17 +1599,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2903,17 +2061,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3372,17 +2519,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3837,16 +2973,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseFeaturestoreServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/client.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/client.py index 10856a55be..593ea860ee 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -136,77 +143,13 @@ def get_transport_class( class GenAiCacheServiceClient(metaclass=GenAiCacheServiceClientMeta): """Service for managing Vertex AI's CachedContent resource.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -416,7 +359,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = GenAiCacheServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -437,35 +380,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = GenAiCacheServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -485,65 +405,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = GenAiCacheServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = GenAiCacheServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -677,13 +538,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - GenAiCacheServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = GenAiCacheServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = GenAiCacheServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=GenAiCacheServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -719,14 +582,14 @@ def __init__( self._transport = cast(GenAiCacheServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or GenAiCacheServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=GenAiCacheServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest.py index 520d7aad8a..bcfcf1d386 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -745,21 +746,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_create_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent, + "_BaseCreateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -899,17 +897,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_delete_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent, + "_BaseDeleteCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1019,17 +1018,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_get_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent, + "_BaseGetCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1169,17 +1169,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_http_options() ) - request, metadata = self._interceptor.pre_list_cached_contents( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListCachedContents, + "_BaseListCachedContents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1327,21 +1328,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_http_options() ) - request, metadata = self._interceptor.pre_update_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent, + "_BaseUpdateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1537,15 +1535,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1679,15 +1678,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1821,19 +1821,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1968,19 +1965,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2116,21 +2110,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2262,17 +2253,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2378,17 +2370,18 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2496,15 +2489,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2638,15 +2632,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2780,15 +2775,16 @@ def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_asyncio.py index 8bafe8a8a6..0c3b340509 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -820,21 +822,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_create_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent, + "_BaseCreateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -978,17 +977,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent, + "_BaseDeleteCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1102,17 +1102,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_get_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent, + "_BaseGetCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1258,17 +1259,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_http_options() ) - request, metadata = await self._interceptor.pre_list_cached_contents( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListCachedContents, + "_BaseListCachedContents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1422,21 +1424,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_update_cached_content( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent, + "_BaseUpdateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1627,17 +1626,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1779,17 +1779,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1932,21 +1933,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2090,21 +2088,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2248,21 +2243,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2400,17 +2392,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2522,17 +2515,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2647,17 +2641,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2799,17 +2794,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2951,17 +2947,18 @@ async def __call__( http_options = ( _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiCacheServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiCacheServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_base.py index 6bb130cdf8..c28b6662b6 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_cache_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.CreateCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseCreateCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteCachedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,43 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.DeleteCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseDeleteCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetCachedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,43 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.GetCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseGetCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListCachedContents: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -271,29 +161,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.ListCachedContentsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseListCachedContents._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateCachedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -302,14 +169,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -321,38 +180,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = gen_ai_cache_service.UpdateCachedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiCacheServiceRestTransport._BaseUpdateCachedContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -371,17 +198,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -400,17 +216,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -494,22 +299,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -600,22 +389,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -691,22 +464,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1117,17 +874,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1582,17 +1328,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2055,17 +1790,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2524,17 +2248,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2989,16 +2702,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseGenAiCacheServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/client.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/client.py index 04e788ab90..ad456d86b0 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -136,77 +143,13 @@ def get_transport_class( class GenAiTuningServiceClient(metaclass=GenAiTuningServiceClientMeta): """A service for creating and managing GenAI Tuning Jobs.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -484,7 +427,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = GenAiTuningServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -505,35 +448,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = GenAiTuningServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -553,65 +473,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = GenAiTuningServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = GenAiTuningServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -745,13 +606,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - GenAiTuningServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = GenAiTuningServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = GenAiTuningServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=GenAiTuningServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -787,14 +650,14 @@ def __init__( self._transport = cast(GenAiTuningServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or GenAiTuningServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=GenAiTuningServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest.py index fde8eb3525..19ee487401 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2962,21 +2963,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob, + "_BaseCancelTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3085,21 +3083,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob, + "_BaseCreateTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3241,15 +3236,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_tuning_job(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob, + "_BaseGetTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3390,17 +3386,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_tuning_jobs( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs, + "_BaseListTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3545,21 +3542,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_http_options() ) - request, metadata = self._interceptor.pre_rebase_tuned_model( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel, + "_BaseRebaseTunedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3746,15 +3740,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3888,15 +3883,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4031,19 +4027,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4179,19 +4172,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4327,21 +4317,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4475,17 +4462,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4591,17 +4579,18 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4710,15 +4699,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4852,15 +4842,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4994,15 +4985,16 @@ def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_asyncio.py index 9b1dbb7b28..023cd9255c 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -814,21 +816,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob, + "_BaseCancelTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -943,21 +942,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob, + "_BaseCreateTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1105,17 +1101,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tuning_job( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob, + "_BaseGetTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1264,17 +1261,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_tuning_jobs( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs, + "_BaseListTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1425,21 +1423,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_http_options() ) - request, metadata = await self._interceptor.pre_rebase_tuned_model( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel, + "_BaseRebaseTunedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3849,17 +3844,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4001,17 +3997,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4154,21 +4151,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4312,21 +4306,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4470,21 +4461,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4622,17 +4610,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4744,17 +4733,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4869,17 +4859,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5021,17 +5012,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5171,17 +5163,18 @@ async def __call__( http_options = ( _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseGenAiTuningServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseGenAiTuningServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_base.py index 89529cf6d1..49b77f8d8c 100644 --- a/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/gen_ai_tuning_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.CancelTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseCancelTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -178,52 +130,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.CreateTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseCreateTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -234,43 +146,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.GetTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseGetTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTuningJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -281,43 +162,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.ListTuningJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseListTuningJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRebaseTunedModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -329,38 +179,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = genai_tuning_service.RebaseTunedModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseGenAiTuningServiceRestTransport._BaseRebaseTunedModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -379,17 +197,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -408,17 +215,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -502,22 +298,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -608,22 +388,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -699,22 +463,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1125,17 +873,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1590,17 +1327,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2063,17 +1789,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2532,17 +2247,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2997,16 +2701,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseGenAiTuningServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/client.py b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/client.py index 159ebcac76..99a6dc0d3c 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class IndexEndpointServiceClient(metaclass=IndexEndpointServiceClientMeta): """A service for managing Vertex AI's IndexEndpoints.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -440,7 +383,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = IndexEndpointServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -461,35 +404,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = IndexEndpointServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -509,65 +429,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = IndexEndpointServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = IndexEndpointServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -701,13 +562,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - IndexEndpointServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = IndexEndpointServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = IndexEndpointServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=IndexEndpointServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -743,14 +606,14 @@ def __init__( self._transport = cast(IndexEndpointServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or IndexEndpointServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=IndexEndpointServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest.py index 5eef05b5ad..634e628186 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3179,21 +3180,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_create_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint, + "_BaseCreateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3336,17 +3334,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_delete_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint, + "_BaseDeleteIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3489,19 +3488,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_http_options() ) - request, metadata = self._interceptor.pre_deploy_index(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeployIndex, + "_BaseDeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3642,17 +3638,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_get_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint, + "_BaseGetIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3795,17 +3792,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_http_options() ) - request, metadata = self._interceptor.pre_list_index_endpoints( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints, + "_BaseListIndexEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3954,21 +3952,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_http_options() ) - request, metadata = self._interceptor.pre_mutate_deployed_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex, + "_BaseMutateDeployedIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4112,19 +4107,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_http_options() ) - request, metadata = self._interceptor.pre_undeploy_index(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex, + "_BaseUndeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4266,21 +4258,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_http_options() ) - request, metadata = self._interceptor.pre_update_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint, + "_BaseUpdateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4508,15 +4497,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4650,15 +4640,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4793,19 +4784,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4941,19 +4929,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5089,21 +5074,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5237,17 +5219,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5353,17 +5336,18 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5472,15 +5456,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5614,15 +5599,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5756,15 +5742,16 @@ def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_asyncio.py index 9119a903d8..1dac894434 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1046,21 +1048,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_create_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint, + "_BaseCreateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1209,17 +1208,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint, + "_BaseDeleteIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1368,21 +1368,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_deploy_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeployIndex, + "_BaseDeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1533,17 +1530,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_get_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint, + "_BaseGetIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1690,17 +1688,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_http_options() ) - request, metadata = await self._interceptor.pre_list_index_endpoints( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints, + "_BaseListIndexEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1853,21 +1852,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_mutate_deployed_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex, + "_BaseMutateDeployedIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2017,21 +2013,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_undeploy_index( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex, + "_BaseUndeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2181,21 +2174,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_http_options() ) - request, metadata = await self._interceptor.pre_update_index_endpoint( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint, + "_BaseUpdateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4636,17 +4626,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4788,17 +4779,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4939,21 +4931,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5095,21 +5084,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5251,21 +5237,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5403,17 +5386,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5525,17 +5509,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5650,17 +5635,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5800,17 +5786,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5950,17 +5937,18 @@ async def __call__( http_options = ( _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexEndpointServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexEndpointServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_base.py index b43f2a9e55..b877b9c75a 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/index_endpoint_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,52 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.CreateIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseCreateIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteIndexEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -176,43 +128,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.DeleteIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseDeleteIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeployIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,52 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.DeployIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseDeployIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetIndexEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -280,43 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.GetIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseGetIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListIndexEndpoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -327,43 +177,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.ListIndexEndpointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseListIndexEndpoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseMutateDeployedIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -375,52 +194,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.MutateDeployedIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseMutateDeployedIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUndeployIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -432,38 +211,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.UndeployIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseUndeployIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateIndexEndpoint: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -472,14 +219,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -491,38 +230,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_endpoint_service.UpdateIndexEndpointRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexEndpointServiceRestTransport._BaseUpdateIndexEndpoint._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -541,17 +248,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -570,17 +266,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -664,22 +349,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -770,22 +439,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -861,22 +514,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1287,17 +924,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1752,17 +1378,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2225,17 +1840,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2694,17 +2298,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3159,16 +2752,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseIndexEndpointServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/index_service/client.py b/google/cloud/aiplatform_v1beta1/services/index_service/client.py index 4ba5f9874f..00a8fba306 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/index_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -139,77 +146,13 @@ class IndexServiceClient(metaclass=IndexServiceClientMeta): resources. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -419,7 +362,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = IndexServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -440,35 +383,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = IndexServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -488,65 +408,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = IndexServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = IndexServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = IndexServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -676,13 +537,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - IndexServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = IndexServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = IndexServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=IndexServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -718,11 +581,14 @@ def __init__( self._transport = cast(IndexServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or IndexServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=IndexServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=IndexServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest.py index 143034a7ab..8b136afaca 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3165,25 +3166,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCreateIndex._get_http_options() ) - request, metadata = self._interceptor.pre_create_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCreateIndex, + "_BaseCreateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3323,19 +3315,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteIndex._get_http_options() ) - request, metadata = self._interceptor.pre_delete_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteIndex, + "_BaseDeleteIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3473,19 +3462,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIndex._get_http_options() ) - request, metadata = self._interceptor.pre_get_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIndex, + "_BaseGetIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3627,25 +3613,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseImportIndex._get_http_options() ) - request, metadata = self._interceptor.pre_import_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseImportIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseImportIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseImportIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseImportIndex, + "_BaseImportIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3784,19 +3761,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListIndexes._get_http_options() ) - request, metadata = self._interceptor.pre_list_indexes(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListIndexes, + "_BaseListIndexes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3939,21 +3913,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_http_options() ) - request, metadata = self._interceptor.pre_remove_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseRemoveDatapoints, + "_BaseRemoveDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4098,25 +4069,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpdateIndex._get_http_options() ) - request, metadata = self._interceptor.pre_update_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpdateIndex, + "_BaseUpdateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4256,21 +4218,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_http_options() ) - request, metadata = self._interceptor.pre_upsert_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpsertDatapoints, + "_BaseUpsertDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4481,19 +4440,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4626,15 +4582,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4768,23 +4725,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4919,23 +4869,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5070,21 +5013,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5215,17 +5155,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5330,17 +5271,18 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5448,17 +5390,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5591,15 +5532,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5732,15 +5674,16 @@ def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseIndexServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_asyncio.py index f176376c12..91cc02032f 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1032,27 +1034,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCreateIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_create_index( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCreateIndex, + "_BaseCreateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1200,21 +1193,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_index( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteIndex, + "_BaseDeleteIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1362,19 +1352,16 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_get_index(request, metadata) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIndex, + "_BaseGetIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1522,27 +1509,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseImportIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_import_index( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseImportIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseImportIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseImportIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseImportIndex, + "_BaseImportIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1689,21 +1667,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListIndexes._get_http_options() ) - request, metadata = await self._interceptor.pre_list_indexes( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseListIndexes._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListIndexes, + "_BaseListIndexes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1852,21 +1827,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_http_options() ) - request, metadata = await self._interceptor.pre_remove_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseRemoveDatapoints, + "_BaseRemoveDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2019,27 +1991,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpdateIndex._get_http_options() ) - request, metadata = await self._interceptor.pre_update_index( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpdateIndex, + "_BaseUpdateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2187,21 +2150,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_http_options() ) - request, metadata = await self._interceptor.pre_upsert_datapoints( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseUpsertDatapoints, + "_BaseUpsertDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4629,21 +4589,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4782,17 +4739,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4934,25 +4892,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5093,25 +5044,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5253,21 +5197,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5406,17 +5347,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5529,17 +5471,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5655,19 +5598,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseIndexServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5806,17 +5748,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5957,17 +5900,18 @@ async def __call__( http_options = ( _BaseIndexServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseIndexServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseIndexServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseIndexServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_base.py index 4be67ceb2d..c4ee97cba5 100644 --- a/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/index_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.CreateIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseCreateIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,43 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.DeleteIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseDeleteIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -222,43 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.GetIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseGetIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -270,52 +160,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.ImportIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseImportIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListIndexes: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -326,43 +176,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.ListIndexesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseListIndexes._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRemoveDatapoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -374,52 +193,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.RemoveDatapointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseRemoveDatapoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateIndex: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -431,52 +210,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.UpdateIndexRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseUpdateIndex._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpsertDatapoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -488,38 +227,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = index_service.UpsertDatapointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseIndexServiceRestTransport._BaseUpsertDatapoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -538,17 +245,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -567,17 +263,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -661,22 +346,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -767,22 +436,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -858,22 +511,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1284,17 +921,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1749,17 +1375,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2222,17 +1837,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2691,17 +2295,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3156,16 +2749,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseIndexServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/job_service/client.py b/google/cloud/aiplatform_v1beta1/services/job_service/client.py index f4c16c2399..ae12d6ce85 100644 --- a/google/cloud/aiplatform_v1beta1/services/job_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/job_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -166,77 +173,13 @@ def get_transport_class( class JobServiceClient(metaclass=JobServiceClientMeta): """A service for creating and managing Vertex AI's jobs.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -801,7 +744,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = JobServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -822,35 +765,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = JobServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -870,65 +790,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = JobServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = JobServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = JobServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -1058,13 +919,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - JobServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = JobServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = JobServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=JobServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -1100,11 +963,14 @@ def __init__( self._transport = cast(JobServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or JobServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=JobServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=JobServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest.py index 242bd1251e..27f6dac7bb 100644 --- a/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -4480,21 +4481,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob, + "_BaseCancelBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4596,21 +4594,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelCustomJob, + "_BaseCancelCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4712,21 +4707,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelDataLabelingJob, + "_BaseCancelDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4830,21 +4822,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob, + "_BaseCancelHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4948,25 +4937,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelNasJob, + "_BaseCancelNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5080,21 +5060,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob, + "_BaseCreateBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5244,21 +5221,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateCustomJob, + "_BaseCreateCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5401,21 +5375,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateDataLabelingJob, + "_BaseCreateDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5563,21 +5534,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob, + "_BaseCreateHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5732,23 +5700,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_create_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob, + "_BaseCreateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5898,25 +5863,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateNasJob, + "_BaseCreateNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6059,17 +6015,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob, + "_BaseDeleteBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6208,17 +6165,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteCustomJob, + "_BaseDeleteCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6357,17 +6315,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob, + "_BaseDeleteDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6508,17 +6467,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob, + "_BaseDeleteHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6663,19 +6623,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_delete_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob, + "_BaseDeleteModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6816,19 +6777,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteNasJob, + "_BaseDeleteNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6971,17 +6929,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetBatchPredictionJob, + "_BaseGetBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7129,19 +7088,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetCustomJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_custom_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetCustomJob, + "_BaseGetCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7282,17 +7238,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetDataLabelingJob, + "_BaseGetDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7437,17 +7394,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_hyperparameter_tuning_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob, + "_BaseGetHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7600,19 +7558,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_get_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob, + "_BaseGetModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7758,19 +7717,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_nas_job(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasJob, + "_BaseGetNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7912,17 +7868,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_http_options() ) - request, metadata = self._interceptor.pre_get_nas_trial_detail( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasTrialDetail, + "_BaseGetNasTrialDetail__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8063,17 +8020,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_batch_prediction_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListBatchPredictionJobs, + "_BaseListBatchPredictionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8215,19 +8173,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListCustomJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_custom_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListCustomJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListCustomJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListCustomJobs, + "_BaseListCustomJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8369,17 +8326,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_data_labeling_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListDataLabelingJobs, + "_BaseListDataLabelingJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8523,17 +8481,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_hyperparameter_tuning_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs, + "_BaseListHyperparameterTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8683,19 +8642,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_http_options() ) - request, metadata = ( self._interceptor.pre_list_model_deployment_monitoring_jobs( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs, + "_BaseListModelDeploymentMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8841,19 +8801,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_nas_jobs(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasJobs, + "_BaseListNasJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8993,17 +8950,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_http_options() ) - request, metadata = self._interceptor.pre_list_nas_trial_details( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasTrialDetails, + "_BaseListNasTrialDetails__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9142,23 +9100,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_pause_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob, + "_BasePauseModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9262,23 +9217,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_resume_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob, + "_BaseResumeModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9390,23 +9342,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_http_options() ) - request, metadata = ( self._interceptor.pre_search_model_deployment_monitoring_stats_anomalies( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies, + "_BaseSearchModelDeploymentMonitoringStatsAnomalies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9561,23 +9510,20 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( self._interceptor.pre_update_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob, + "_BaseUpdateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10053,19 +9999,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10198,19 +10141,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10344,25 +10284,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10497,25 +10428,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10650,21 +10572,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10795,17 +10714,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10910,17 +10830,18 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11028,19 +10949,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11173,17 +11091,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseJobServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11316,19 +11233,16 @@ def __call__( http_options = ( _BaseJobServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_asyncio.py index e3134c9ab0..868dae4971 100644 --- a/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -2485,21 +2487,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob, + "_BaseCancelBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2607,21 +2606,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelCustomJob, + "_BaseCancelCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2731,21 +2727,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelDataLabelingJob, + "_BaseCancelDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2857,23 +2850,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_cancel_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob, + "_BaseCancelHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2981,27 +2971,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelNasJob, + "_BaseCancelNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3122,21 +3103,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob, + "_BaseCreateBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3294,21 +3272,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateCustomJob, + "_BaseCreateCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3459,21 +3434,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateDataLabelingJob, + "_BaseCreateDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3631,23 +3603,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob, + "_BaseCreateHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3808,23 +3777,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob, + "_BaseCreateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3982,27 +3948,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCreateNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCreateNasJob, + "_BaseCreateNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4152,17 +4109,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob, + "_BaseDeleteBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4311,17 +4269,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_custom_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteCustomJob, + "_BaseDeleteCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4470,17 +4429,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob, + "_BaseDeleteDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4633,19 +4593,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob, + "_BaseDeleteHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4798,19 +4759,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob, + "_BaseDeleteModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4961,21 +4923,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteNasJob, + "_BaseDeleteNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5126,17 +5085,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_batch_prediction_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetBatchPredictionJob, + "_BaseGetBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5294,21 +5254,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetCustomJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_custom_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetCustomJob, + "_BaseGetCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5455,17 +5412,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_data_labeling_job( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetDataLabelingJob, + "_BaseGetDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5619,19 +5577,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_hyperparameter_tuning_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob, + "_BaseGetHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5788,19 +5747,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob, + "_BaseGetModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5954,21 +5914,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_nas_job( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetNasJob._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasJob, + "_BaseGetNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6116,17 +6073,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_http_options() ) - request, metadata = await self._interceptor.pre_get_nas_trial_detail( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetNasTrialDetail, + "_BaseGetNasTrialDetail__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6276,17 +6234,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_batch_prediction_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListBatchPredictionJobs, + "_BaseListBatchPredictionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6436,19 +6395,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListCustomJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_custom_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListCustomJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListCustomJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListCustomJobs, + "_BaseListCustomJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6596,17 +6554,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_data_labeling_jobs( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListDataLabelingJobs, + "_BaseListDataLabelingJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6760,19 +6719,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_hyperparameter_tuning_jobs( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs, + "_BaseListHyperparameterTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6928,19 +6888,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_model_deployment_monitoring_jobs( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs, + "_BaseListModelDeploymentMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7094,21 +7055,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_nas_jobs( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListNasJobs._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasJobs, + "_BaseListNasJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7254,17 +7212,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_http_options() ) - request, metadata = await self._interceptor.pre_list_nas_trial_details( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListNasTrialDetails, + "_BaseListNasTrialDetails__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7413,23 +7372,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_pause_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob, + "_BasePauseModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7541,23 +7497,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_resume_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob, + "_BaseResumeModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7675,23 +7628,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_http_options() ) - request, metadata = ( await self._interceptor.pre_search_model_deployment_monitoring_stats_anomalies( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies, + "_BaseSearchModelDeploymentMonitoringStatsAnomalies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7854,23 +7804,20 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_model_deployment_monitoring_job( request, metadata ) ) - transcoded_request = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob, + "_BaseUpdateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10511,21 +10458,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10664,21 +10608,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10818,27 +10759,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10979,27 +10911,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11140,21 +11063,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11293,17 +11213,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11416,17 +11337,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseJobServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11542,21 +11464,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11695,19 +11614,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseJobServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11846,21 +11764,18 @@ async def __call__( http_options = ( _BaseJobServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseJobServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseJobServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_base.py index f028b3abbc..b2645511a3 100644 --- a/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/job_service/transports/rest_base.py @@ -120,14 +120,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -139,52 +131,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -196,52 +148,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -253,52 +165,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -310,52 +182,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -367,52 +199,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CancelNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCancelNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateBatchPredictionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -424,52 +216,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -481,52 +233,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -538,52 +250,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -595,52 +267,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -652,54 +284,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -711,52 +301,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.CreateNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseCreateNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteBatchPredictionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -767,43 +317,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -814,43 +333,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -861,43 +349,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -908,43 +365,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -955,45 +381,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1004,43 +397,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.DeleteNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseDeleteNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetBatchPredictionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1051,43 +413,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetBatchPredictionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetBatchPredictionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetCustomJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1098,43 +429,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetCustomJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetCustomJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetDataLabelingJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1145,43 +445,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetDataLabelingJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetDataLabelingJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetHyperparameterTuningJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1192,43 +461,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetHyperparameterTuningJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetHyperparameterTuningJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1239,43 +477,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetModelDeploymentMonitoringJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNasJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1286,43 +493,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetNasJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetNasJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNasTrialDetail: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1333,43 +509,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.GetNasTrialDetailRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseGetNasTrialDetail._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListBatchPredictionJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1380,43 +525,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListBatchPredictionJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListBatchPredictionJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListCustomJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1427,43 +541,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListCustomJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListCustomJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListDataLabelingJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1474,43 +557,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListDataLabelingJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListDataLabelingJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListHyperparameterTuningJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1521,43 +573,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListHyperparameterTuningJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListHyperparameterTuningJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelDeploymentMonitoringJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1568,45 +589,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListModelDeploymentMonitoringJobsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListModelDeploymentMonitoringJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNasJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1617,43 +605,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListNasJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListNasJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNasTrialDetails: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1664,43 +621,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ListNasTrialDetailsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseListNasTrialDetails._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePauseModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1712,54 +638,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.PauseModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BasePauseModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseResumeModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1771,54 +655,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.ResumeModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseResumeModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchModelDeploymentMonitoringStatsAnomalies: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1830,42 +672,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - job_service.SearchModelDeploymentMonitoringStatsAnomaliesRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseSearchModelDeploymentMonitoringStatsAnomalies._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateModelDeploymentMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1874,14 +680,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1893,40 +691,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = job_service.UpdateModelDeploymentMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseJobServiceRestTransport._BaseUpdateModelDeploymentMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1945,17 +709,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1974,17 +727,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2068,22 +810,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2174,22 +900,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2265,22 +975,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2691,17 +1385,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3156,17 +1839,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3629,17 +2301,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4098,17 +2759,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4563,16 +3213,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseJobServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/client.py b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/client.py index 2baf281705..0f5a9faae6 100644 --- a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -128,77 +135,13 @@ def get_transport_class( class LlmUtilityServiceClient(metaclass=LlmUtilityServiceClientMeta): """Service for LLM related utility functions.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -386,7 +329,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = LlmUtilityServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -407,35 +350,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = LlmUtilityServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -455,65 +375,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = LlmUtilityServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = LlmUtilityServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -647,13 +508,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - LlmUtilityServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = LlmUtilityServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = LlmUtilityServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=LlmUtilityServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -689,14 +552,14 @@ def __init__( self._transport = cast(LlmUtilityServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or LlmUtilityServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=LlmUtilityServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest.py index 4f5a56b5bd..8d28edc150 100644 --- a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -547,19 +548,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_http_options() ) - request, metadata = self._interceptor.pre_compute_tokens(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseComputeTokens, + "_BaseComputeTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -713,15 +711,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -855,15 +854,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -997,19 +997,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1144,19 +1141,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1292,21 +1286,18 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1438,17 +1429,18 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1554,17 +1546,18 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1672,15 +1665,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1814,15 +1808,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1956,15 +1951,16 @@ def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_asyncio.py index df8538043f..e541ae62d9 100644 --- a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -602,21 +604,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_http_options() ) - request, metadata = await self._interceptor.pre_compute_tokens( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseComputeTokens, + "_BaseComputeTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -777,17 +776,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -929,17 +929,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1082,21 +1083,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1240,21 +1238,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1398,21 +1393,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1550,17 +1542,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1672,17 +1665,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1797,17 +1791,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1949,17 +1944,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2101,17 +2097,18 @@ async def __call__( http_options = ( _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseLlmUtilityServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseLlmUtilityServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_base.py index 16f1f48fc0..b39ba97006 100644 --- a/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/llm_utility_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -133,38 +125,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = llm_utility_service.ComputeTokensRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseLlmUtilityServiceRestTransport._BaseComputeTokens._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -183,17 +143,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -212,17 +161,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -306,22 +244,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -412,22 +334,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -503,22 +409,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -929,17 +819,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1394,17 +1273,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1867,17 +1735,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2336,17 +2193,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2801,16 +2647,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseLlmUtilityServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/match_service/client.py b/google/cloud/aiplatform_v1beta1/services/match_service/client.py index 2be3cc7390..90936aa39d 100644 --- a/google/cloud/aiplatform_v1beta1/services/match_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/match_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -128,77 +135,13 @@ class MatchServiceClient(metaclass=MatchServiceClientMeta): similarity search at scale. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -386,7 +329,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = MatchServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -407,35 +350,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = MatchServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -455,65 +375,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = MatchServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = MatchServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = MatchServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -643,13 +504,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - MatchServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = MatchServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = MatchServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=MatchServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -685,11 +548,14 @@ def __init__( self._transport = cast(MatchServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or MatchServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=MatchServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=MatchServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest.py index b0dc53de28..b8c5f7a523 100644 --- a/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -605,19 +606,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseFindNeighbors._get_http_options() ) - request, metadata = self._interceptor.pre_find_neighbors(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseFindNeighbors, + "_BaseFindNeighbors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -761,21 +759,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_http_options() ) - request, metadata = self._interceptor.pre_read_index_datapoints( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseReadIndexDatapoints, + "_BaseReadIndexDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -939,19 +934,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1084,15 +1076,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1226,23 +1219,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1377,23 +1363,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1528,21 +1507,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1673,17 +1649,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1788,17 +1765,18 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1906,17 +1884,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2049,15 +2026,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2190,15 +2168,16 @@ def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseMatchServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_asyncio.py index 37fec19d72..880d97a323 100644 --- a/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -665,21 +667,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseFindNeighbors._get_http_options() ) - request, metadata = await self._interceptor.pre_find_neighbors( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseFindNeighbors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseFindNeighbors, + "_BaseFindNeighbors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -832,21 +831,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_http_options() ) - request, metadata = await self._interceptor.pre_read_index_datapoints( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseReadIndexDatapoints, + "_BaseReadIndexDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1014,21 +1010,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1167,17 +1160,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1319,25 +1313,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1478,25 +1465,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1638,21 +1618,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1791,17 +1768,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1914,17 +1892,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2040,19 +2019,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseMatchServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2191,17 +2169,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2342,17 +2321,18 @@ async def __call__( http_options = ( _BaseMatchServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseMatchServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMatchServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMatchServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_base.py index f14c7ad72f..aadb7bd2fe 100644 --- a/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/match_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,52 +110,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = match_service.FindNeighborsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMatchServiceRestTransport._BaseFindNeighbors._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadIndexDatapoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,38 +127,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = match_service.ReadIndexDatapointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMatchServiceRestTransport._BaseReadIndexDatapoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -225,17 +145,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -254,17 +163,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -348,22 +246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -454,22 +336,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -545,22 +411,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -971,17 +821,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1436,17 +1275,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1909,17 +1737,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2378,17 +2195,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2843,16 +2649,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseMatchServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/async_client.py b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/async_client.py index 7b9949cabb..e03b17d2f2 100644 --- a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/async_client.py +++ b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/async_client.py @@ -345,12 +345,8 @@ async def sample_create_memory(): client = aiplatform_v1beta1.MemoryBankServiceAsyncClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.CreateMemoryRequest( parent="parent_value", - memory=memory, ) # Make the request @@ -609,11 +605,7 @@ async def sample_update_memory(): client = aiplatform_v1beta1.MemoryBankServiceAsyncClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.UpdateMemoryRequest( - memory=memory, ) # Make the request diff --git a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/client.py b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/client.py index 5f193b2d59..90b4bab7c0 100644 --- a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -135,77 +142,13 @@ def get_transport_class( class MemoryBankServiceClient(metaclass=MemoryBankServiceClientMeta): """A service for managing memories for LLM applications.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -441,7 +384,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = MemoryBankServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -462,35 +405,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = MemoryBankServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -510,65 +430,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = MemoryBankServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = MemoryBankServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = MemoryBankServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = MemoryBankServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -702,13 +563,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - MemoryBankServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = MemoryBankServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = MemoryBankServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=MemoryBankServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -744,14 +607,14 @@ def __init__( self._transport = cast(MemoryBankServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or MemoryBankServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=MemoryBankServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=MemoryBankServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=MemoryBankServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: @@ -866,12 +729,8 @@ def sample_create_memory(): client = aiplatform_v1beta1.MemoryBankServiceClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.CreateMemoryRequest( parent="parent_value", - memory=memory, ) # Make the request @@ -1124,11 +983,7 @@ def sample_update_memory(): client = aiplatform_v1beta1.MemoryBankServiceClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.UpdateMemoryRequest( - memory=memory, ) # Make the request diff --git a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest.py index 7dd27b152e..047c52a98a 100644 --- a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3114,19 +3115,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_http_options() ) - request, metadata = self._interceptor.pre_create_memory(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseCreateMemory, + "_BaseCreateMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3266,15 +3264,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseDeleteMemory._get_http_options() ) - request, metadata = self._interceptor.pre_delete_memory(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseDeleteMemory._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseDeleteMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseDeleteMemory, + "_BaseDeleteMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3415,21 +3414,18 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_http_options() ) - request, metadata = self._interceptor.pre_generate_memories( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGenerateMemories, + "_BaseGenerateMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3566,15 +3562,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetMemory._get_http_options() ) - request, metadata = self._interceptor.pre_get_memory(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetMemory._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetMemory, + "_BaseGetMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3714,15 +3711,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseListMemories._get_http_options() ) - request, metadata = self._interceptor.pre_list_memories(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseListMemories._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseListMemories._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseListMemories, + "_BaseListMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3866,21 +3864,18 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_http_options() ) - request, metadata = self._interceptor.pre_retrieve_memories( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories, + "_BaseRetrieveMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4025,19 +4020,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_http_options() ) - request, metadata = self._interceptor.pre_update_memory(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseUpdateMemory, + "_BaseUpdateMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4240,15 +4232,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4382,15 +4375,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4524,19 +4518,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4671,19 +4662,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4819,21 +4807,18 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4965,17 +4950,18 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5081,17 +5067,18 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5199,15 +5186,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5341,15 +5329,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5483,15 +5472,16 @@ def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_asyncio.py index 20a9eb13fd..a8a5cafc90 100644 --- a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -977,21 +979,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_http_options() ) - request, metadata = await self._interceptor.pre_create_memory( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseCreateMemory, + "_BaseCreateMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1142,17 +1141,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseDeleteMemory._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_memory( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseDeleteMemory._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseDeleteMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseDeleteMemory, + "_BaseDeleteMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1303,21 +1303,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_http_options() ) - request, metadata = await self._interceptor.pre_generate_memories( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGenerateMemories, + "_BaseGenerateMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1463,17 +1460,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetMemory._get_http_options() ) - request, metadata = await self._interceptor.pre_get_memory( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetMemory._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetMemory, + "_BaseGetMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1622,17 +1620,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseListMemories._get_http_options() ) - request, metadata = await self._interceptor.pre_list_memories( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseListMemories._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseListMemories._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseListMemories, + "_BaseListMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1784,21 +1783,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_http_options() ) - request, metadata = await self._interceptor.pre_retrieve_memories( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories, + "_BaseRetrieveMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1950,21 +1946,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_http_options() ) - request, metadata = await self._interceptor.pre_update_memory( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseUpdateMemory, + "_BaseUpdateMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4389,17 +4382,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4541,17 +4535,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4694,21 +4689,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4852,21 +4844,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5010,21 +4999,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5162,17 +5148,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5284,17 +5271,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5409,17 +5397,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5561,17 +5550,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5713,17 +5703,18 @@ async def __call__( http_options = ( _BaseMemoryBankServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseMemoryBankServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMemoryBankServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMemoryBankServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_base.py index 01acae70d3..4dccab87a9 100644 --- a/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/memory_bank_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -124,52 +116,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = memory_bank_service.CreateMemoryRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMemoryBankServiceRestTransport._BaseCreateMemory._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteMemory: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -184,43 +136,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = memory_bank_service.DeleteMemoryRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMemoryBankServiceRestTransport._BaseDeleteMemory._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGenerateMemories: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -237,52 +158,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = memory_bank_service.GenerateMemoriesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMemoryBankServiceRestTransport._BaseGenerateMemories._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetMemory: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -297,43 +178,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = memory_bank_service.GetMemoryRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMemoryBankServiceRestTransport._BaseGetMemory._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListMemories: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -348,43 +198,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = memory_bank_service.ListMemoriesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMemoryBankServiceRestTransport._BaseListMemories._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRetrieveMemories: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -401,52 +220,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = memory_bank_service.RetrieveMemoriesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMemoryBankServiceRestTransport._BaseRetrieveMemories._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateMemory: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -463,38 +242,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = memory_bank_service.UpdateMemoryRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMemoryBankServiceRestTransport._BaseUpdateMemory._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -513,17 +260,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -542,17 +278,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -636,22 +361,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -742,22 +451,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -833,22 +526,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1259,17 +936,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1724,17 +1390,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2197,17 +1852,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2666,17 +2310,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3131,16 +2764,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseMemoryBankServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/metadata_service/client.py b/google/cloud/aiplatform_v1beta1/services/metadata_service/client.py index b1397673d5..8efda971f6 100644 --- a/google/cloud/aiplatform_v1beta1/services/metadata_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/metadata_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -148,77 +155,13 @@ def get_transport_class( class MetadataServiceClient(metaclass=MetadataServiceClientMeta): """Service for reading and writing metadata entries.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -502,7 +445,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = MetadataServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -523,35 +466,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = MetadataServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -571,65 +491,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = MetadataServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = MetadataServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -761,13 +622,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - MetadataServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = MetadataServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = MetadataServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=MetadataServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -803,14 +666,14 @@ def __init__( self._transport = cast(MetadataServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or MetadataServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=MetadataServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=MetadataServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest.py index 8be057cd90..cc74d85788 100644 --- a/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -4554,23 +4555,20 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_http_options() ) - request, metadata = ( self._interceptor.pre_add_context_artifacts_and_executions( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions, + "_BaseAddContextArtifactsAndExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4717,21 +4715,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_http_options() ) - request, metadata = self._interceptor.pre_add_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextChildren, + "_BaseAddContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4876,21 +4871,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_http_options() ) - request, metadata = self._interceptor.pre_add_execution_events( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddExecutionEvents, + "_BaseAddExecutionEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5032,19 +5024,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_create_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateArtifact, + "_BaseCreateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5184,19 +5173,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateContext._get_http_options() ) - request, metadata = self._interceptor.pre_create_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateContext, + "_BaseCreateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5336,21 +5322,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateExecution._get_http_options() ) - request, metadata = self._interceptor.pre_create_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateExecution, + "_BaseCreateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5491,21 +5474,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_http_options() ) - request, metadata = self._interceptor.pre_create_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema, + "_BaseCreateMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5651,21 +5631,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_http_options() ) - request, metadata = self._interceptor.pre_create_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataStore, + "_BaseCreateMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5805,15 +5782,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_delete_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteArtifact, + "_BaseDeleteArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5952,15 +5930,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteContext._get_http_options() ) - request, metadata = self._interceptor.pre_delete_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteContext, + "_BaseDeleteContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6099,17 +6078,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_http_options() ) - request, metadata = self._interceptor.pre_delete_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteExecution, + "_BaseDeleteExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6249,17 +6229,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_http_options() ) - request, metadata = self._interceptor.pre_delete_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore, + "_BaseDeleteMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6395,15 +6376,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_get_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetArtifact, + "_BaseGetArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6541,15 +6523,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetContext._get_http_options() ) - request, metadata = self._interceptor.pre_get_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetContext, + "_BaseGetContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6687,15 +6670,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetExecution._get_http_options() ) - request, metadata = self._interceptor.pre_get_execution(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetExecution, + "_BaseGetExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6834,17 +6818,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_http_options() ) - request, metadata = self._interceptor.pre_get_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataSchema, + "_BaseGetMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6985,17 +6970,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_http_options() ) - request, metadata = self._interceptor.pre_get_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataStore, + "_BaseGetMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7135,15 +7121,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListArtifacts._get_http_options() ) - request, metadata = self._interceptor.pre_list_artifacts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListArtifacts, + "_BaseListArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7285,15 +7272,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListContexts._get_http_options() ) - request, metadata = self._interceptor.pre_list_contexts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListContexts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListContexts, + "_BaseListContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7435,15 +7423,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListExecutions._get_http_options() ) - request, metadata = self._interceptor.pre_list_executions(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListExecutions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListExecutions, + "_BaseListExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7586,17 +7575,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_http_options() ) - request, metadata = self._interceptor.pre_list_metadata_schemas( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataSchemas, + "_BaseListMetadataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7739,17 +7729,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_http_options() ) - request, metadata = self._interceptor.pre_list_metadata_stores( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataStores, + "_BaseListMetadataStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7893,19 +7884,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_http_options() ) - request, metadata = self._interceptor.pre_purge_artifacts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeArtifacts, + "_BasePurgeArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8046,19 +8034,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeContexts._get_http_options() ) - request, metadata = self._interceptor.pre_purge_contexts(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeContexts, + "_BasePurgeContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8199,21 +8184,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_http_options() ) - request, metadata = self._interceptor.pre_purge_executions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeExecutions, + "_BasePurgeExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8355,17 +8337,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_http_options() ) - request, metadata = self._interceptor.pre_query_artifact_lineage_subgraph( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph, + "_BaseQueryArtifactLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8512,17 +8495,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_http_options() ) - request, metadata = self._interceptor.pre_query_context_lineage_subgraph( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph, + "_BaseQueryContextLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8671,19 +8655,20 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_http_options() ) - request, metadata = ( self._interceptor.pre_query_execution_inputs_and_outputs( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs, + "_BaseQueryExecutionInputsAndOutputs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8829,21 +8814,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_http_options() ) - request, metadata = self._interceptor.pre_remove_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseRemoveContextChildren, + "_BaseRemoveContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8987,19 +8969,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_http_options() ) - request, metadata = self._interceptor.pre_update_artifact(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateArtifact, + "_BaseUpdateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9139,19 +9118,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateContext._get_http_options() ) - request, metadata = self._interceptor.pre_update_context(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateContext, + "_BaseUpdateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9291,21 +9267,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_http_options() ) - request, metadata = self._interceptor.pre_update_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateExecution, + "_BaseUpdateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9747,15 +9720,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9888,15 +9862,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10030,19 +10005,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10177,19 +10149,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10325,21 +10294,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10470,17 +10436,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10585,17 +10552,18 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10703,15 +10671,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10844,15 +10813,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10985,15 +10955,16 @@ def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_asyncio.py index 281705c0c5..f134c3c39c 100644 --- a/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -2549,23 +2551,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_http_options() ) - request, metadata = ( await self._interceptor.pre_add_context_artifacts_and_executions( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions, + "_BaseAddContextArtifactsAndExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2720,21 +2719,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_http_options() ) - request, metadata = await self._interceptor.pre_add_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddContextChildren, + "_BaseAddContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2885,21 +2881,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_http_options() ) - request, metadata = await self._interceptor.pre_add_execution_events( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseAddExecutionEvents, + "_BaseAddExecutionEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3048,21 +3041,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_create_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateArtifact, + "_BaseCreateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3211,21 +3201,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateContext._get_http_options() ) - request, metadata = await self._interceptor.pre_create_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateContext, + "_BaseCreateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3374,21 +3361,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_create_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateExecution, + "_BaseCreateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3537,21 +3521,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_http_options() ) - request, metadata = await self._interceptor.pre_create_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema, + "_BaseCreateMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3703,21 +3684,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_http_options() ) - request, metadata = await self._interceptor.pre_create_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCreateMetadataStore, + "_BaseCreateMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3866,17 +3844,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteArtifact, + "_BaseDeleteArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4026,17 +4005,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteContext._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteContext, + "_BaseDeleteContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4186,17 +4166,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteExecution, + "_BaseDeleteExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4346,17 +4327,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore, + "_BaseDeleteMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4500,17 +4482,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_get_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetArtifact, + "_BaseGetArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4656,17 +4639,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetContext._get_http_options() ) - request, metadata = await self._interceptor.pre_get_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetContext._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetContext, + "_BaseGetContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4813,17 +4797,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_get_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetExecution._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetExecution, + "_BaseGetExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4970,17 +4955,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_http_options() ) - request, metadata = await self._interceptor.pre_get_metadata_schema( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataSchema, + "_BaseGetMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5128,17 +5114,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_http_options() ) - request, metadata = await self._interceptor.pre_get_metadata_store( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetMetadataStore, + "_BaseGetMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5287,17 +5274,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListArtifacts._get_http_options() ) - request, metadata = await self._interceptor.pre_list_artifacts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListArtifacts, + "_BaseListArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5448,17 +5436,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_list_contexts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListContexts._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListContexts, + "_BaseListContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5609,17 +5598,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListExecutions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_executions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListExecutions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListExecutions, + "_BaseListExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5770,17 +5760,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_http_options() ) - request, metadata = await self._interceptor.pre_list_metadata_schemas( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataSchemas, + "_BaseListMetadataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5929,17 +5920,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_http_options() ) - request, metadata = await self._interceptor.pre_list_metadata_stores( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListMetadataStores, + "_BaseListMetadataStores__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6090,21 +6082,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_http_options() ) - request, metadata = await self._interceptor.pre_purge_artifacts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeArtifacts, + "_BasePurgeArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6256,21 +6245,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_purge_contexts( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeContexts, + "_BasePurgeContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6422,21 +6408,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_http_options() ) - request, metadata = await self._interceptor.pre_purge_executions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BasePurgeExecutions, + "_BasePurgeExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6590,19 +6573,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_http_options() ) - request, metadata = ( await self._interceptor.pre_query_artifact_lineage_subgraph( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph, + "_BaseQueryArtifactLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6755,19 +6739,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_http_options() ) - request, metadata = ( await self._interceptor.pre_query_context_lineage_subgraph( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph, + "_BaseQueryContextLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6922,19 +6907,20 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_query_execution_inputs_and_outputs( request, metadata ) ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs, + "_BaseQueryExecutionInputsAndOutputs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7086,21 +7072,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_http_options() ) - request, metadata = await self._interceptor.pre_remove_context_children( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseRemoveContextChildren, + "_BaseRemoveContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7251,21 +7234,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_http_options() ) - request, metadata = await self._interceptor.pre_update_artifact( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateArtifact, + "_BaseUpdateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7414,21 +7394,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateContext._get_http_options() ) - request, metadata = await self._interceptor.pre_update_context( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateContext._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateContext, + "_BaseUpdateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7577,21 +7554,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_http_options() ) - request, metadata = await self._interceptor.pre_update_execution( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseUpdateExecution, + "_BaseUpdateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10202,17 +10176,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10354,17 +10329,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10507,21 +10483,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10665,21 +10638,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10823,21 +10793,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10975,17 +10942,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11099,17 +11067,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11226,17 +11195,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11378,17 +11348,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11530,17 +11501,18 @@ async def __call__( http_options = ( _BaseMetadataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMetadataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMetadataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_base.py index ed5d6a3cfd..5c12cf8e38 100644 --- a/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/metadata_service/transports/rest_base.py @@ -109,14 +109,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -128,54 +120,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.AddContextArtifactsAndExecutionsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseAddContextArtifactsAndExecutions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAddContextChildren: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -187,52 +137,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.AddContextChildrenRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseAddContextChildren._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAddExecutionEvents: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -244,52 +154,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.AddExecutionEventsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseAddExecutionEvents._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -301,52 +171,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -358,52 +188,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -415,52 +205,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateMetadataSchema: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -472,52 +222,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateMetadataSchemaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateMetadataSchema._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateMetadataStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -529,52 +239,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.CreateMetadataStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseCreateMetadataStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -585,43 +255,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -632,43 +271,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -679,43 +287,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteMetadataStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -726,43 +303,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.DeleteMetadataStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseDeleteMetadataStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -773,43 +319,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -820,43 +335,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -867,43 +351,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetMetadataSchema: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -914,43 +367,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetMetadataSchemaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetMetadataSchema._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetMetadataStore: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -961,43 +383,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.GetMetadataStoreRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseGetMetadataStore._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListArtifacts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1008,43 +399,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListArtifactsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListArtifacts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1055,43 +415,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListExecutions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1102,43 +431,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListExecutionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListExecutions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListMetadataSchemas: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1149,43 +447,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListMetadataSchemasRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListMetadataSchemas._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListMetadataStores: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1196,43 +463,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.ListMetadataStoresRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseListMetadataStores._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePurgeArtifacts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1244,52 +480,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.PurgeArtifactsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BasePurgeArtifacts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePurgeContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1301,52 +497,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.PurgeContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BasePurgeContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePurgeExecutions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1358,52 +514,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.PurgeExecutionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BasePurgeExecutions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryArtifactLineageSubgraph: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1414,45 +530,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.QueryArtifactLineageSubgraphRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseQueryArtifactLineageSubgraph._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryContextLineageSubgraph: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1463,43 +546,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.QueryContextLineageSubgraphRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseQueryContextLineageSubgraph._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryExecutionInputsAndOutputs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1510,45 +562,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.QueryExecutionInputsAndOutputsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseQueryExecutionInputsAndOutputs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRemoveContextChildren: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1560,52 +579,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.RemoveContextChildrenRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseRemoveContextChildren._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateArtifact: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1617,52 +596,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.UpdateArtifactRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseUpdateArtifact._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateContext: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1674,52 +613,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.UpdateContextRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseUpdateContext._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateExecution: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1731,38 +630,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = metadata_service.UpdateExecutionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMetadataServiceRestTransport._BaseUpdateExecution._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1781,17 +648,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1810,17 +666,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1904,22 +749,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2010,22 +839,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2101,22 +914,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2527,17 +1324,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2992,17 +1778,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3465,17 +2240,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3934,17 +2698,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4399,16 +3152,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseMetadataServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/migration_service/client.py b/google/cloud/aiplatform_v1beta1/services/migration_service/client.py index dd0127c6b1..19cc38d45c 100644 --- a/google/cloud/aiplatform_v1beta1/services/migration_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/migration_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -133,77 +140,13 @@ class MigrationServiceClient(metaclass=MigrationServiceClientMeta): datalabeling.googleapis.com and ml.googleapis.com to Vertex AI. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -518,7 +461,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = MigrationServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -539,35 +482,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = MigrationServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -587,65 +507,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = MigrationServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = MigrationServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -777,13 +638,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - MigrationServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = MigrationServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = MigrationServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=MigrationServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -819,14 +682,14 @@ def __init__( self._transport = cast(MigrationServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or MigrationServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=MigrationServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=MigrationServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest.py index fb20670c00..36fc980c16 100644 --- a/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2834,21 +2835,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_http_options() ) - request, metadata = self._interceptor.pre_batch_migrate_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseBatchMigrateResources, + "_BaseBatchMigrateResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2992,21 +2990,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_http_options() ) - request, metadata = self._interceptor.pre_search_migratable_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSearchMigratableResources, + "_BaseSearchMigratableResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3174,15 +3169,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3315,15 +3311,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3457,19 +3454,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3604,19 +3598,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3752,21 +3743,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3898,17 +3886,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4014,17 +4003,18 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4132,15 +4122,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4273,15 +4264,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4414,15 +4406,16 @@ def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_asyncio.py index 63f438591e..872e0345df 100644 --- a/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -671,21 +673,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_migrate_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseBatchMigrateResources, + "_BaseBatchMigrateResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -837,21 +836,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_http_options() ) - request, metadata = await self._interceptor.pre_search_migratable_resources( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSearchMigratableResources, + "_BaseSearchMigratableResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3247,17 +3243,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3399,17 +3396,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3552,21 +3550,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3710,21 +3705,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3868,21 +3860,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4020,17 +4009,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4144,17 +4134,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4271,17 +4262,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4423,17 +4415,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4575,17 +4568,18 @@ async def __call__( http_options = ( _BaseMigrationServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseMigrationServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseMigrationServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_base.py index 4929865cd5..1ca45475c1 100644 --- a/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/migration_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,52 +110,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = migration_service.BatchMigrateResourcesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMigrationServiceRestTransport._BaseBatchMigrateResources._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchMigratableResources: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,38 +127,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = migration_service.SearchMigratableResourcesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseMigrationServiceRestTransport._BaseSearchMigratableResources._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -225,17 +145,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -254,17 +163,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -348,22 +246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -454,22 +336,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -545,22 +411,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -971,17 +821,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1436,17 +1275,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1909,17 +1737,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2378,17 +2195,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2843,16 +2649,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseMigrationServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/model_garden_service/client.py b/google/cloud/aiplatform_v1beta1/services/model_garden_service/client.py index 696c1c46d8..f003309f58 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_garden_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/model_garden_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -132,77 +139,13 @@ def get_transport_class( class ModelGardenServiceClient(metaclass=ModelGardenServiceClientMeta): """The interface of Model Garden Service.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -451,7 +394,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ModelGardenServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -472,35 +415,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ModelGardenServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -520,65 +440,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ModelGardenServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ModelGardenServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -712,13 +573,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ModelGardenServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ModelGardenServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ModelGardenServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ModelGardenServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -754,14 +617,14 @@ def __init__( self._transport = cast(ModelGardenServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ModelGardenServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ModelGardenServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest.py index 8c43bfe8bc..e9721f1c79 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3124,21 +3125,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_http_options() ) - request, metadata = self._interceptor.pre_accept_publisher_model_eula( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula, + "_BaseAcceptPublisherModelEula__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3290,23 +3288,20 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_http_options() ) - request, metadata = ( self._interceptor.pre_check_publisher_model_eula_acceptance( request, metadata ) ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance, + "_BaseCheckPublisherModelEulaAcceptance__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3455,23 +3450,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeploy._get_http_options() ) - request, metadata = self._interceptor.pre_deploy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeploy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeploy, + "_BaseDeploy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3613,21 +3601,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_http_options() ) - request, metadata = self._interceptor.pre_deploy_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel, + "_BaseDeployPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3771,21 +3756,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_http_options() ) - request, metadata = self._interceptor.pre_export_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseExportPublisherModel, + "_BaseExportPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3925,17 +3907,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_http_options() ) - request, metadata = self._interceptor.pre_get_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetPublisherModel, + "_BaseGetPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4076,17 +4059,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListPublisherModels._get_http_options() ) - request, metadata = self._interceptor.pre_list_publisher_models( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListPublisherModels._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListPublisherModels._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListPublisherModels, + "_BaseListPublisherModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4304,15 +4288,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4446,15 +4431,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4589,19 +4575,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4737,19 +4720,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4885,21 +4865,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5033,17 +5010,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5149,17 +5127,18 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5268,15 +5247,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5410,15 +5390,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5552,15 +5533,16 @@ def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_asyncio.py index 52f9c1bafa..2380ab0006 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -986,21 +988,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_http_options() ) - request, metadata = await self._interceptor.pre_accept_publisher_model_eula( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula, + "_BaseAcceptPublisherModelEula__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1158,23 +1157,20 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_http_options() ) - request, metadata = ( await self._interceptor.pre_check_publisher_model_eula_acceptance( request, metadata ) ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance, + "_BaseCheckPublisherModelEulaAcceptance__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1332,23 +1328,16 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeploy._get_http_options() ) - request, metadata = await self._interceptor.pre_deploy(request, metadata) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeploy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeploy, + "_BaseDeploy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1498,21 +1487,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_http_options() ) - request, metadata = await self._interceptor.pre_deploy_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel, + "_BaseDeployPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1662,21 +1648,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_http_options() ) - request, metadata = await self._interceptor.pre_export_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseExportPublisherModel, + "_BaseExportPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1822,17 +1805,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_http_options() ) - request, metadata = await self._interceptor.pre_get_publisher_model( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetPublisherModel, + "_BaseGetPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1979,17 +1963,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListPublisherModels._get_http_options() ) - request, metadata = await self._interceptor.pre_list_publisher_models( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListPublisherModels._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListPublisherModels._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListPublisherModels, + "_BaseListPublisherModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4422,17 +4407,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4574,17 +4560,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4727,21 +4714,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4885,21 +4869,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5043,21 +5024,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5195,17 +5173,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5317,17 +5296,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5442,17 +5422,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5594,17 +5575,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5744,17 +5726,18 @@ async def __call__( http_options = ( _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelGardenServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelGardenServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_base.py index 15aac8ffc3..8bd6291152 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/model_garden_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,54 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.AcceptPublisherModelEulaRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseAcceptPublisherModelEula._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCheckPublisherModelEulaAcceptance: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -178,56 +128,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - model_garden_service.CheckPublisherModelEulaAcceptanceRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseCheckPublisherModelEulaAcceptance._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeploy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -239,52 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.DeployRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseDeploy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeployPublisherModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -296,52 +162,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.DeployPublisherModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseDeployPublisherModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportPublisherModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -353,52 +179,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.ExportPublisherModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseExportPublisherModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetPublisherModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -409,43 +195,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.GetPublisherModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseGetPublisherModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListPublisherModels: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -456,29 +211,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_garden_service.ListPublisherModelsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelGardenServiceRestTransport._BaseListPublisherModels._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -497,17 +229,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -526,17 +247,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -620,22 +330,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -726,22 +420,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -817,22 +495,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1243,17 +905,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1708,17 +1359,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2181,17 +1821,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2650,17 +2279,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3115,16 +2733,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseModelGardenServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/client.py b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/client.py index e5e0580d87..665f97b692 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -149,77 +156,13 @@ class ModelMonitoringServiceClient(metaclass=ModelMonitoringServiceClientMeta): resources. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -563,7 +506,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ModelMonitoringServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -584,35 +527,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ModelMonitoringServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -632,67 +552,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ModelMonitoringServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ModelMonitoringServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - ModelMonitoringServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ModelMonitoringServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -826,13 +685,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ModelMonitoringServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ModelMonitoringServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ModelMonitoringServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ModelMonitoringServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -868,14 +729,14 @@ def __init__( self._transport = cast(ModelMonitoringServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ModelMonitoringServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ModelMonitoringServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ModelMonitoringServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ModelMonitoringServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest.py index e8f2a51b3f..387d587b71 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3367,21 +3368,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_create_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor, + "_BaseCreateModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3526,21 +3524,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_model_monitoring_job( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob, + "_BaseCreateModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3685,17 +3680,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_delete_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor, + "_BaseDeleteModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3838,17 +3834,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_model_monitoring_job( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob, + "_BaseDeleteModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3992,17 +3989,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_get_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor, + "_BaseGetModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4146,17 +4144,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_model_monitoring_job( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob, + "_BaseGetModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4300,17 +4299,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_monitoring_jobs( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs, + "_BaseListModelMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4453,17 +4453,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_monitors( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors, + "_BaseListModelMonitors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4614,21 +4615,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_http_options() ) - request, metadata = self._interceptor.pre_search_model_monitoring_alerts( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts, + "_BaseSearchModelMonitoringAlerts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4780,21 +4778,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_http_options() ) - request, metadata = self._interceptor.pre_search_model_monitoring_stats( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats, + "_BaseSearchModelMonitoringStats__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4944,21 +4939,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_http_options() ) - request, metadata = self._interceptor.pre_update_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor, + "_BaseUpdateModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5217,15 +5209,16 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5359,15 +5352,16 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5502,19 +5496,16 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5650,19 +5641,16 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5798,21 +5786,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5946,17 +5931,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6064,17 +6050,18 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6185,15 +6172,16 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6327,15 +6315,16 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6471,15 +6460,16 @@ def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_asyncio.py index 4bad22450d..14a5d139ab 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1251,21 +1253,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_create_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor, + "_BaseCreateModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1418,21 +1417,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_model_monitoring_job( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob, + "_BaseCreateModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1585,17 +1581,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor, + "_BaseDeleteModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1746,17 +1743,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_model_monitoring_job( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob, + "_BaseDeleteModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1910,17 +1908,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor, + "_BaseGetModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2070,17 +2069,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model_monitoring_job( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob, + "_BaseGetModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2234,17 +2234,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_model_monitoring_jobs( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs, + "_BaseListModelMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2395,17 +2396,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors._get_http_options() ) - request, metadata = await self._interceptor.pre_list_model_monitors( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors, + "_BaseListModelMonitors__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2560,23 +2562,20 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_http_options() ) - request, metadata = ( await self._interceptor.pre_search_model_monitoring_alerts( request, metadata ) ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts, + "_BaseSearchModelMonitoringAlerts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2734,23 +2733,20 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_http_options() ) - request, metadata = ( await self._interceptor.pre_search_model_monitoring_stats( request, metadata ) ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats, + "_BaseSearchModelMonitoringStats__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2906,21 +2902,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_http_options() ) - request, metadata = await self._interceptor.pre_update_model_monitor( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor, + "_BaseUpdateModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5388,17 +5381,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5538,17 +5532,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5689,21 +5684,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5845,21 +5837,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6001,21 +5990,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6153,17 +6139,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6275,17 +6262,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6400,17 +6388,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6550,17 +6539,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6700,17 +6690,18 @@ async def __call__( http_options = ( _BaseModelMonitoringServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseModelMonitoringServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelMonitoringServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelMonitoringServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_base.py index ca04020071..55429dc65f 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/model_monitoring_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -123,52 +115,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.CreateModelMonitorRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateModelMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -180,54 +132,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.CreateModelMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseCreateModelMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModelMonitor: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -238,43 +148,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.DeleteModelMonitorRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModelMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -285,45 +164,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.DeleteModelMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseDeleteModelMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelMonitor: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -334,43 +180,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.GetModelMonitorRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelMonitoringJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -381,45 +196,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.GetModelMonitoringJobRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseGetModelMonitoringJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelMonitoringJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -430,45 +212,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.ListModelMonitoringJobsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseListModelMonitoringJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelMonitors: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -479,43 +228,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.ListModelMonitorsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseListModelMonitors._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchModelMonitoringAlerts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -527,54 +245,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.SearchModelMonitoringAlertsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringAlerts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSearchModelMonitoringStats: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -586,40 +262,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.SearchModelMonitoringStatsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseSearchModelMonitoringStats._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateModelMonitor: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -628,14 +270,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -647,38 +281,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_monitoring_service.UpdateModelMonitorRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelMonitoringServiceRestTransport._BaseUpdateModelMonitor._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -697,17 +299,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -726,17 +317,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -820,22 +400,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -926,22 +490,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1017,22 +565,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1443,17 +975,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1908,17 +1429,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2381,17 +1891,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2850,17 +2349,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3315,16 +2803,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseModelMonitoringServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/model_service/client.py b/google/cloud/aiplatform_v1beta1/services/model_service/client.py index 4eeeb0fe78..e7c43adb1f 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/model_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -144,77 +151,13 @@ def get_transport_class( class ModelServiceClient(metaclass=ModelServiceClientMeta): """A service for managing Vertex AI's machine learning Models.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -518,7 +461,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ModelServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -539,35 +482,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ModelServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -587,65 +507,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ModelServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ModelServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ModelServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -775,13 +636,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ModelServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ModelServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ModelServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ModelServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -817,11 +680,14 @@ def __init__( self._transport = cast(ModelServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = self._api_endpoint or ModelServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ModelServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ModelServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest.py index c75d5458b8..9100898034 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3866,23 +3867,20 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_import_evaluated_annotations( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations, + "_BaseBatchImportEvaluatedAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4032,23 +4030,20 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_import_model_evaluation_slices( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices, + "_BaseBatchImportModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4197,23 +4192,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCopyModel._get_http_options() ) - request, metadata = self._interceptor.pre_copy_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_transcoded_request( - http_options, request - ) - ) - - body = _BaseModelServiceRestTransport._BaseCopyModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCopyModel, + "_BaseCopyModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4353,19 +4341,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModel._get_http_options() ) - request, metadata = self._interceptor.pre_delete_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModel, + "_BaseDeleteModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4504,17 +4489,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_http_options() ) - request, metadata = self._interceptor.pre_delete_model_version( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModelVersion, + "_BaseDeleteModelVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4654,25 +4640,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseExportModel._get_http_options() ) - request, metadata = self._interceptor.pre_export_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseExportModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseExportModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseExportModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseExportModel, + "_BaseExportModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4807,19 +4784,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModel._get_http_options() ) - request, metadata = self._interceptor.pre_get_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModel, + "_BaseGetModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4961,17 +4935,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_http_options() ) - request, metadata = self._interceptor.pre_get_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluation, + "_BaseGetModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5117,17 +5092,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_http_options() ) - request, metadata = self._interceptor.pre_get_model_evaluation_slice( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice, + "_BaseGetModelEvaluationSlice__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5272,21 +5248,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_http_options() ) - request, metadata = self._interceptor.pre_import_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseImportModelEvaluation, + "_BaseImportModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5429,17 +5402,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluations._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_evaluations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluations, + "_BaseListModelEvaluations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5583,17 +5557,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_evaluation_slices( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluationSlices, + "_BaseListModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5739,19 +5714,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModels._get_http_options() ) - request, metadata = self._interceptor.pre_list_models(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseListModels._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseListModels._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModels, + "_BaseListModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5895,17 +5867,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_version_checkpoints( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints, + "_BaseListModelVersionCheckpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6053,17 +6026,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersions._get_http_options() ) - request, metadata = self._interceptor.pre_list_model_versions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersions, + "_BaseListModelVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6204,21 +6178,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_http_options() ) - request, metadata = self._interceptor.pre_merge_version_aliases( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseMergeVersionAliases, + "_BaseMergeVersionAliases__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6360,19 +6331,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseRecommendSpec._get_http_options() ) - request, metadata = self._interceptor.pre_recommend_spec(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseRecommendSpec._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseRecommendSpec._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseRecommendSpec._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseRecommendSpec, + "_BaseRecommendSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6519,21 +6487,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_http_options() ) - request, metadata = self._interceptor.pre_update_explanation_dataset( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateExplanationDataset, + "_BaseUpdateExplanationDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6673,25 +6638,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateModel._get_http_options() ) - request, metadata = self._interceptor.pre_update_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateModel, + "_BaseUpdateModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6834,25 +6790,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUploadModel._get_http_options() ) - request, metadata = self._interceptor.pre_upload_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUploadModel, + "_BaseUploadModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7179,19 +7126,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7324,15 +7268,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7466,23 +7411,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7617,23 +7555,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7768,21 +7699,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7913,17 +7841,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8028,17 +7957,18 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8146,17 +8076,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8289,15 +8218,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8430,15 +8360,16 @@ def __call__( http_options = ( _BaseModelServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseModelServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_asyncio.py index 80ccd8edaf..792fdf620d 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1795,23 +1797,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_import_evaluated_annotations( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations, + "_BaseBatchImportEvaluatedAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1969,23 +1968,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_import_model_evaluation_slices( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices, + "_BaseBatchImportModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2142,25 +2138,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCopyModel._get_http_options() ) - request, metadata = await self._interceptor.pre_copy_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_transcoded_request( - http_options, request - ) - ) - - body = _BaseModelServiceRestTransport._BaseCopyModel._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseCopyModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCopyModel, + "_BaseCopyModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2308,21 +2297,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModel._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseDeleteModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModel, + "_BaseDeleteModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2470,17 +2456,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_model_version( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteModelVersion, + "_BaseDeleteModelVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2630,27 +2617,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseExportModel._get_http_options() ) - request, metadata = await self._interceptor.pre_export_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseExportModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseExportModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseExportModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseExportModel, + "_BaseExportModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2795,19 +2773,16 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModel._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model(request, metadata) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetModel._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModel, + "_BaseGetModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2956,17 +2931,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluation, + "_BaseGetModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3120,17 +3096,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_http_options() ) - request, metadata = await self._interceptor.pre_get_model_evaluation_slice( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice, + "_BaseGetModelEvaluationSlice__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3284,21 +3261,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_http_options() ) - request, metadata = await self._interceptor.pre_import_model_evaluation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseImportModelEvaluation, + "_BaseImportModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3450,17 +3424,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_model_evaluations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluations, + "_BaseListModelEvaluations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3610,19 +3585,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_model_evaluation_slices( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelEvaluationSlices, + "_BaseListModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3774,21 +3750,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModels._get_http_options() ) - request, metadata = await self._interceptor.pre_list_models( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseListModels._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseListModels._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModels, + "_BaseListModels__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3938,19 +3911,20 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_model_version_checkpoints( request, metadata ) ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints, + "_BaseListModelVersionCheckpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4102,17 +4076,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListModelVersions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_model_versions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListModelVersions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListModelVersions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListModelVersions, + "_BaseListModelVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4262,21 +4237,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_http_options() ) - request, metadata = await self._interceptor.pre_merge_version_aliases( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseMergeVersionAliases, + "_BaseMergeVersionAliases__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4426,21 +4398,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseRecommendSpec._get_http_options() ) - request, metadata = await self._interceptor.pre_recommend_spec( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseRecommendSpec._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseRecommendSpec._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseRecommendSpec._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseRecommendSpec, + "_BaseRecommendSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4595,21 +4564,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_http_options() ) - request, metadata = await self._interceptor.pre_update_explanation_dataset( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateExplanationDataset, + "_BaseUpdateExplanationDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4757,27 +4723,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUpdateModel._get_http_options() ) - request, metadata = await self._interceptor.pre_update_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUpdateModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUpdateModel, + "_BaseUpdateModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4926,27 +4883,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseUploadModel._get_http_options() ) - request, metadata = await self._interceptor.pre_upload_model( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseUploadModel._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseUploadModel, + "_BaseUploadModel__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7466,21 +7414,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7619,17 +7564,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7771,25 +7717,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7930,25 +7869,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8090,21 +8022,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8243,17 +8172,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8366,17 +8296,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8492,19 +8423,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseModelServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8643,17 +8573,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8794,17 +8725,18 @@ async def __call__( http_options = ( _BaseModelServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseModelServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseModelServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseModelServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_base.py index c738047d55..d714cfb2f5 100644 --- a/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/model_service/transports/rest_base.py @@ -106,14 +106,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -125,54 +117,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.BatchImportEvaluatedAnnotationsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseBatchImportEvaluatedAnnotations._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchImportModelEvaluationSlices: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -184,54 +134,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.BatchImportModelEvaluationSlicesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseBatchImportModelEvaluationSlices._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCopyModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -243,52 +151,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.CopyModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseCopyModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -299,43 +167,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.DeleteModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseDeleteModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteModelVersion: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -346,43 +183,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.DeleteModelVersionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseDeleteModelVersion._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -394,52 +200,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ExportModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseExportModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -450,43 +216,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.GetModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseGetModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelEvaluation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -497,43 +232,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.GetModelEvaluationRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseGetModelEvaluation._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetModelEvaluationSlice: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -544,43 +248,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.GetModelEvaluationSliceRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseGetModelEvaluationSlice._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportModelEvaluation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -592,52 +265,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ImportModelEvaluationRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseImportModelEvaluation._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelEvaluations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -648,43 +281,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelEvaluationsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelEvaluations._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelEvaluationSlices: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -695,43 +297,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelEvaluationSlicesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelEvaluationSlices._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModels: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -742,43 +313,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModels._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelVersionCheckpoints: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -789,43 +329,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelVersionCheckpointsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelVersionCheckpoints._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListModelVersions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -836,43 +345,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.ListModelVersionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseListModelVersions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseMergeVersionAliases: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -884,52 +362,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.MergeVersionAliasesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseMergeVersionAliases._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRecommendSpec: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -941,52 +379,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.RecommendSpecRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseRecommendSpec._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateExplanationDataset: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -998,38 +396,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.UpdateExplanationDatasetRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseUpdateExplanationDataset._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1038,14 +404,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1057,52 +415,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.UpdateModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseUpdateModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUploadModel: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1114,38 +432,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = model_service.UploadModelRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseModelServiceRestTransport._BaseUploadModel._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1164,17 +450,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1193,17 +468,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1287,22 +551,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1393,22 +641,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1484,22 +716,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1910,17 +1126,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2375,17 +1580,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2848,17 +2042,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3317,17 +2500,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3782,16 +2954,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseModelServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/notebook_service/client.py b/google/cloud/aiplatform_v1beta1/services/notebook_service/client.py index a218fde063..e87514c5e8 100644 --- a/google/cloud/aiplatform_v1beta1/services/notebook_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/notebook_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -154,77 +161,13 @@ class NotebookServiceClient(metaclass=NotebookServiceClientMeta): Workbench). """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -541,7 +484,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = NotebookServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -562,35 +505,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = NotebookServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -610,65 +530,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = NotebookServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = NotebookServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -800,13 +661,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - NotebookServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = NotebookServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = NotebookServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=NotebookServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -842,14 +705,14 @@ def __init__( self._transport = cast(NotebookServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or NotebookServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=NotebookServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=NotebookServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest.py index 885d2fdd3a..cf351b7720 100644 --- a/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3651,21 +3652,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_assign_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime, + "_BaseAssignNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3810,21 +3808,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob, + "_BaseCreateNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3971,21 +3966,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_create_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate, + "_BaseCreateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4129,17 +4121,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob, + "_BaseDeleteNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4283,17 +4276,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_delete_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime, + "_BaseDeleteNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4436,17 +4430,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_delete_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate, + "_BaseDeleteNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4588,17 +4583,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob, + "_BaseGetNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4748,17 +4744,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_get_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime, + "_BaseGetNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4905,17 +4902,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_get_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate, + "_BaseGetNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5063,17 +5061,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_notebook_execution_jobs( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs, + "_BaseListNotebookExecutionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5220,17 +5219,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_http_options() ) - request, metadata = self._interceptor.pre_list_notebook_runtimes( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes, + "_BaseListNotebookRuntimes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5374,17 +5374,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_http_options() ) - request, metadata = self._interceptor.pre_list_notebook_runtime_templates( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates, + "_BaseListNotebookRuntimeTemplates__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5533,21 +5534,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_start_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime, + "_BaseStartNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5689,21 +5687,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_stop_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime, + "_BaseStopNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5848,21 +5843,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = self._interceptor.pre_update_notebook_runtime_template( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate, + "_BaseUpdateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6010,21 +6002,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_http_options() ) - request, metadata = self._interceptor.pre_upgrade_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime, + "_BaseUpgradeNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6333,15 +6322,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6474,15 +6464,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6616,19 +6607,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6763,19 +6751,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6911,21 +6896,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7056,17 +7038,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7171,17 +7154,18 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7289,15 +7273,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7430,15 +7415,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7571,15 +7557,16 @@ def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_asyncio.py index 0e3905360b..d452aa1968 100644 --- a/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1558,21 +1560,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_assign_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime, + "_BaseAssignNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1725,23 +1724,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_notebook_execution_job( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob, + "_BaseCreateNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1896,23 +1892,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate, + "_BaseCreateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2064,19 +2057,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_notebook_execution_job( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob, + "_BaseDeleteNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2226,17 +2220,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime, + "_BaseDeleteNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2389,19 +2384,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate, + "_BaseDeleteNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2551,17 +2547,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_notebook_execution_job( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob, + "_BaseGetNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2717,17 +2714,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_get_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime, + "_BaseGetNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2880,19 +2878,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate, + "_BaseGetNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3044,19 +3043,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_notebook_execution_jobs( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs, + "_BaseListNotebookExecutionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3209,17 +3209,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_http_options() ) - request, metadata = await self._interceptor.pre_list_notebook_runtimes( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes, + "_BaseListNotebookRuntimes__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3371,19 +3372,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_notebook_runtime_templates( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates, + "_BaseListNotebookRuntimeTemplates__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3538,21 +3540,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_start_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime, + "_BaseStartNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3702,21 +3701,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_stop_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime, + "_BaseStopNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3871,23 +3867,20 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_notebook_runtime_template( request, metadata ) ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate, + "_BaseUpdateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4041,21 +4034,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_http_options() ) - request, metadata = await self._interceptor.pre_upgrade_notebook_runtime( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime, + "_BaseUpgradeNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6565,17 +6555,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6717,17 +6708,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6870,21 +6862,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7028,21 +7017,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7186,21 +7172,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7338,17 +7321,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7462,17 +7446,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7589,17 +7574,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7741,17 +7727,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7893,17 +7880,18 @@ async def __call__( http_options = ( _BaseNotebookServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseNotebookServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseNotebookServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_base.py index 5f0cf651a5..174a39f062 100644 --- a/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/notebook_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,52 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.AssignNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseAssignNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateNotebookExecutionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,52 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.CreateNotebookExecutionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseCreateNotebookExecutionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -234,54 +146,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.CreateNotebookRuntimeTemplateRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseCreateNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNotebookExecutionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -292,43 +162,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.DeleteNotebookExecutionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseDeleteNotebookExecutionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -339,43 +178,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.DeleteNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -386,45 +194,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.DeleteNotebookRuntimeTemplateRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseDeleteNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNotebookExecutionJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -435,43 +210,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.GetNotebookExecutionJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseGetNotebookExecutionJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -482,43 +226,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.GetNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseGetNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -529,43 +242,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.GetNotebookRuntimeTemplateRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseGetNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNotebookExecutionJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -576,43 +258,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.ListNotebookExecutionJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseListNotebookExecutionJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNotebookRuntimes: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -623,43 +274,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.ListNotebookRuntimesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseListNotebookRuntimes._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListNotebookRuntimeTemplates: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -670,45 +290,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.ListNotebookRuntimeTemplatesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseListNotebookRuntimeTemplates._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStartNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -720,52 +307,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.StartNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseStartNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStopNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -777,38 +324,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.StopNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseStopNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateNotebookRuntimeTemplate: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -817,14 +332,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -836,54 +343,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.UpdateNotebookRuntimeTemplateRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseUpdateNotebookRuntimeTemplate._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpgradeNotebookRuntime: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -895,38 +360,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = notebook_service.UpgradeNotebookRuntimeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseNotebookServiceRestTransport._BaseUpgradeNotebookRuntime._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -945,17 +378,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -974,17 +396,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1068,22 +479,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1174,22 +569,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1265,22 +644,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1691,17 +1054,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2156,17 +1508,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2629,17 +1970,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3098,17 +2428,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3563,16 +2882,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseNotebookServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/client.py b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/client.py index f7ce4ec966..65fd4498f8 100644 --- a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -140,77 +147,13 @@ class OnlineEvaluatorServiceClient(metaclass=OnlineEvaluatorServiceClientMeta): OnlineEvaluators. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -398,7 +341,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = OnlineEvaluatorServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -419,35 +362,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = OnlineEvaluatorServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -467,67 +387,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = OnlineEvaluatorServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = OnlineEvaluatorServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - OnlineEvaluatorServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = OnlineEvaluatorServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -661,13 +520,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - OnlineEvaluatorServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = OnlineEvaluatorServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = OnlineEvaluatorServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=OnlineEvaluatorServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -703,14 +564,14 @@ def __init__( self._transport = cast(OnlineEvaluatorServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or OnlineEvaluatorServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=OnlineEvaluatorServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=OnlineEvaluatorServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=OnlineEvaluatorServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest.py index 49a5997357..752cec8746 100644 --- a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3122,21 +3123,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_http_options() ) - request, metadata = self._interceptor.pre_activate_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator, + "_BaseActivateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3278,21 +3276,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_http_options() ) - request, metadata = self._interceptor.pre_create_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator, + "_BaseCreateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3433,17 +3428,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator._get_http_options() ) - request, metadata = self._interceptor.pre_delete_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator, + "_BaseDeleteOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3582,17 +3578,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator._get_http_options() ) - request, metadata = self._interceptor.pre_get_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator, + "_BaseGetOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3737,17 +3734,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators._get_http_options() ) - request, metadata = self._interceptor.pre_list_online_evaluators( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators, + "_BaseListOnlineEvaluators__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3896,21 +3894,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_http_options() ) - request, metadata = self._interceptor.pre_suspend_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator, + "_BaseSuspendOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4052,21 +4047,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_http_options() ) - request, metadata = self._interceptor.pre_update_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator, + "_BaseUpdateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4283,15 +4275,16 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4425,15 +4418,16 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4568,19 +4562,16 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4716,19 +4707,16 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4864,21 +4852,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5012,17 +4997,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5130,17 +5116,18 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5251,15 +5238,16 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5393,15 +5381,16 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5537,15 +5526,16 @@ def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_asyncio.py index 4291f4ee79..90d5122954 100644 --- a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -988,21 +990,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_http_options() ) - request, metadata = await self._interceptor.pre_activate_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator, + "_BaseActivateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1156,21 +1155,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_http_options() ) - request, metadata = await self._interceptor.pre_create_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator, + "_BaseCreateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1323,17 +1319,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator, + "_BaseDeleteOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1482,17 +1479,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator._get_http_options() ) - request, metadata = await self._interceptor.pre_get_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator, + "_BaseGetOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1641,17 +1639,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators._get_http_options() ) - request, metadata = await self._interceptor.pre_list_online_evaluators( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators, + "_BaseListOnlineEvaluators__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1806,21 +1805,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_http_options() ) - request, metadata = await self._interceptor.pre_suspend_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator, + "_BaseSuspendOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1974,21 +1970,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_http_options() ) - request, metadata = await self._interceptor.pre_update_online_evaluator( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator, + "_BaseUpdateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4426,17 +4419,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4576,17 +4570,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4727,21 +4722,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4883,21 +4875,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5039,21 +5028,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5191,17 +5177,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5313,17 +5300,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5438,17 +5426,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5588,17 +5577,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5738,17 +5728,18 @@ async def __call__( http_options = ( _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseOnlineEvaluatorServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_base.py index 2a44d92288..4c2435f1df 100644 --- a/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/online_evaluator_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,54 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = online_evaluator_service.ActivateOnlineEvaluatorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseOnlineEvaluatorServiceRestTransport._BaseActivateOnlineEvaluator._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateOnlineEvaluator: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -178,54 +128,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = online_evaluator_service.CreateOnlineEvaluatorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseOnlineEvaluatorServiceRestTransport._BaseCreateOnlineEvaluator._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteOnlineEvaluator: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -236,45 +144,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = online_evaluator_service.DeleteOnlineEvaluatorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseOnlineEvaluatorServiceRestTransport._BaseDeleteOnlineEvaluator._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetOnlineEvaluator: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -285,43 +160,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = online_evaluator_service.GetOnlineEvaluatorRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseOnlineEvaluatorServiceRestTransport._BaseGetOnlineEvaluator._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListOnlineEvaluators: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -332,45 +176,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = online_evaluator_service.ListOnlineEvaluatorsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseOnlineEvaluatorServiceRestTransport._BaseListOnlineEvaluators._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSuspendOnlineEvaluator: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -382,54 +193,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = online_evaluator_service.SuspendOnlineEvaluatorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseOnlineEvaluatorServiceRestTransport._BaseSuspendOnlineEvaluator._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateOnlineEvaluator: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -441,40 +210,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = online_evaluator_service.UpdateOnlineEvaluatorRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseOnlineEvaluatorServiceRestTransport._BaseUpdateOnlineEvaluator._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -493,17 +228,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -522,17 +246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -616,22 +329,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -722,22 +419,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -813,22 +494,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1239,17 +904,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1704,17 +1358,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2177,17 +1820,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2646,17 +2278,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3111,16 +2732,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseOnlineEvaluatorServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/client.py b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/client.py index 5d5446b962..4e22ac9b2f 100644 --- a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -145,77 +152,13 @@ class PersistentResourceServiceClient(metaclass=PersistentResourceServiceClientM PersistentResource. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -488,7 +431,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = PersistentResourceServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -509,35 +452,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = PersistentResourceServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -557,67 +477,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = PersistentResourceServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = PersistentResourceServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -751,15 +610,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - PersistentResourceServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( PersistentResourceServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = PersistentResourceServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=PersistentResourceServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -795,14 +656,14 @@ def __init__( self._transport = cast(PersistentResourceServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or PersistentResourceServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=PersistentResourceServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest.py index 94cf361eec..5b5171f670 100644 --- a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3070,21 +3071,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_create_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource, + "_BaseCreatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3228,17 +3226,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_delete_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource, + "_BaseDeletePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3380,17 +3379,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_get_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource, + "_BaseGetPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3535,17 +3535,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_http_options() ) - request, metadata = self._interceptor.pre_list_persistent_resources( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListPersistentResources, + "_BaseListPersistentResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3695,21 +3696,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_reboot_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource, + "_BaseRebootPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3854,21 +3852,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_http_options() ) - request, metadata = self._interceptor.pre_update_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource, + "_BaseUpdatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4074,15 +4069,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4218,15 +4214,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4363,19 +4360,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4513,19 +4507,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4663,21 +4654,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4809,17 +4797,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4927,17 +4916,18 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5048,15 +5038,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5192,15 +5183,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5336,15 +5328,16 @@ def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_asyncio.py index 44a193e939..c2091e808c 100644 --- a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -929,21 +931,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_create_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource, + "_BaseCreatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1097,17 +1096,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource, + "_BaseDeletePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1261,17 +1261,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_get_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource, + "_BaseGetPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1424,17 +1425,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_http_options() ) - request, metadata = await self._interceptor.pre_list_persistent_resources( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListPersistentResources, + "_BaseListPersistentResources__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1592,21 +1594,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_reboot_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource, + "_BaseRebootPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1761,21 +1760,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_http_options() ) - request, metadata = await self._interceptor.pre_update_persistent_resource( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource, + "_BaseUpdatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4204,17 +4200,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4354,17 +4351,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4505,21 +4503,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4661,21 +4656,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4819,21 +4811,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4971,17 +4960,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5093,17 +5083,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5218,17 +5209,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5368,17 +5360,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5518,17 +5511,18 @@ async def __call__( http_options = ( _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePersistentResourceServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePersistentResourceServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_base.py index 98cab41b79..da623b9747 100644 --- a/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/persistent_resource_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER "persistentResourceId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,54 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.CreatePersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseCreatePersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeletePersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -179,45 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.DeletePersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseDeletePersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetPersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -228,45 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.GetPersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseGetPersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListPersistentResources: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -277,45 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.ListPersistentResourcesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseListPersistentResources._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRebootPersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -327,40 +178,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.RebootPersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseRebootPersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdatePersistentResource: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -369,14 +186,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -388,40 +197,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = persistent_resource_service.UpdatePersistentResourceRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePersistentResourceServiceRestTransport._BaseUpdatePersistentResource._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -440,17 +215,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -469,17 +233,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -563,22 +316,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -669,22 +406,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -760,22 +481,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1186,17 +891,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1651,17 +1345,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2124,17 +1807,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2593,17 +2265,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3058,16 +2719,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BasePersistentResourceServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/pipeline_service/client.py b/google/cloud/aiplatform_v1beta1/services/pipeline_service/client.py index 33303d8bc9..0626d2bbd8 100644 --- a/google/cloud/aiplatform_v1beta1/services/pipeline_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/pipeline_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -149,77 +156,13 @@ class PipelineServiceClient(metaclass=PipelineServiceClientMeta): Pipelines). """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -608,7 +551,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = PipelineServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -629,35 +572,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = PipelineServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -677,65 +597,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = PipelineServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = PipelineServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -867,13 +728,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - PipelineServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = PipelineServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = PipelineServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=PipelineServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -909,14 +772,14 @@ def __init__( self._transport = cast(PipelineServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or PipelineServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=PipelineServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=PipelineServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest.py index 25646bc167..4d14743d87 100644 --- a/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3344,21 +3345,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_http_options() ) - request, metadata = self._interceptor.pre_batch_cancel_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs, + "_BaseBatchCancelPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3503,21 +3501,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_http_options() ) - request, metadata = self._interceptor.pre_batch_delete_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs, + "_BaseBatchDeletePipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3654,21 +3649,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelPipelineJob, + "_BaseCancelPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3771,21 +3763,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline, + "_BaseCancelTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3896,21 +3885,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_create_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreatePipelineJob, + "_BaseCreatePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4057,21 +4043,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_create_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline, + "_BaseCreateTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4218,17 +4201,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_delete_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeletePipelineJob, + "_BaseDeletePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4368,17 +4352,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_delete_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline, + "_BaseDeleteTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4518,17 +4503,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_http_options() ) - request, metadata = self._interceptor.pre_get_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetPipelineJob, + "_BaseGetPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4673,17 +4659,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_http_options() ) - request, metadata = self._interceptor.pre_get_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetTrainingPipeline, + "_BaseGetTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4825,17 +4812,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_http_options() ) - request, metadata = self._interceptor.pre_list_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListPipelineJobs, + "_BaseListPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4978,17 +4966,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_http_options() ) - request, metadata = self._interceptor.pre_list_training_pipelines( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListTrainingPipelines, + "_BaseListTrainingPipelines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5250,15 +5239,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5391,15 +5381,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5533,19 +5524,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5680,19 +5668,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5828,21 +5813,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5973,17 +5955,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6088,17 +6071,18 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6206,15 +6190,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6347,15 +6332,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6488,15 +6474,16 @@ def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BasePipelineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_asyncio.py index cc7b353c43..280d8aa239 100644 --- a/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1231,21 +1233,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_cancel_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs, + "_BaseBatchCancelPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1398,21 +1397,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_delete_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs, + "_BaseBatchDeletePipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1557,21 +1553,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelPipelineJob, + "_BaseCancelPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1680,21 +1673,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline, + "_BaseCancelTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1809,21 +1799,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_create_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreatePipelineJob, + "_BaseCreatePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1976,21 +1963,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_create_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline, + "_BaseCreateTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2143,17 +2127,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeletePipelineJob, + "_BaseDeletePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2301,17 +2286,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline, + "_BaseDeleteTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2460,17 +2446,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_http_options() ) - request, metadata = await self._interceptor.pre_get_pipeline_job( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetPipelineJob, + "_BaseGetPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2623,17 +2610,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_http_options() ) - request, metadata = await self._interceptor.pre_get_training_pipeline( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetTrainingPipeline, + "_BaseGetTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2782,17 +2770,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_http_options() ) - request, metadata = await self._interceptor.pre_list_pipeline_jobs( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListPipelineJobs, + "_BaseListPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2943,17 +2932,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_http_options() ) - request, metadata = await self._interceptor.pre_list_training_pipelines( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListTrainingPipelines, + "_BaseListTrainingPipelines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5422,17 +5412,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5574,17 +5565,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5727,21 +5719,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5885,21 +5874,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6043,21 +6029,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6195,17 +6178,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6319,17 +6303,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6446,17 +6431,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6598,17 +6584,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6750,17 +6737,18 @@ async def __call__( http_options = ( _BasePipelineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BasePipelineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePipelineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePipelineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_base.py index 6c97b67f6e..683e938194 100644 --- a/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/pipeline_service/transports/rest_base.py @@ -106,14 +106,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -125,52 +117,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.BatchCancelPipelineJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseBatchCancelPipelineJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchDeletePipelineJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -182,52 +134,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.BatchDeletePipelineJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseBatchDeletePipelineJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelPipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -239,52 +151,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CancelPipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCancelPipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -296,52 +168,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CancelTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCancelTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreatePipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -353,52 +185,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CreatePipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCreatePipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -410,52 +202,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.CreateTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseCreateTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeletePipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -466,43 +218,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.DeletePipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseDeletePipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -513,43 +234,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.DeleteTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseDeleteTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetPipelineJob: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -560,43 +250,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.GetPipelineJobRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseGetPipelineJob._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTrainingPipeline: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -607,43 +266,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.GetTrainingPipelineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseGetTrainingPipeline._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListPipelineJobs: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -654,43 +282,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.ListPipelineJobsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseListPipelineJobs._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTrainingPipelines: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -701,29 +298,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = pipeline_service.ListTrainingPipelinesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePipelineServiceRestTransport._BaseListTrainingPipelines._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -742,17 +316,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -771,17 +334,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -865,22 +417,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -971,22 +507,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1062,22 +582,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1488,17 +992,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1953,17 +1446,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2426,17 +1908,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2895,17 +2366,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3360,16 +2820,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BasePipelineServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/prediction_service/client.py b/google/cloud/aiplatform_v1beta1/services/prediction_service/client.py index c944fd3ff0..3619136906 100644 --- a/google/cloud/aiplatform_v1beta1/services/prediction_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/prediction_service/client.py @@ -41,6 +41,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class PredictionServiceClient(metaclass=PredictionServiceClientMeta): """A service for online predictions and explanations.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -484,7 +427,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = PredictionServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -505,35 +448,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = PredictionServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -553,65 +473,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = PredictionServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = PredictionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -745,13 +606,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - PredictionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = PredictionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = PredictionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=PredictionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -787,14 +650,14 @@ def __init__( self._transport = cast(PredictionServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or PredictionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=PredictionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=PredictionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest.py index 3138b04336..cd1de37e05 100644 --- a/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -1238,21 +1239,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseChatCompletions._get_http_options() ) - request, metadata = self._interceptor.pre_chat_completions( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseChatCompletions._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseChatCompletions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseChatCompletions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseChatCompletions, + "_BaseChatCompletions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1386,19 +1384,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseCountTokens._get_http_options() ) - request, metadata = self._interceptor.pre_count_tokens(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseCountTokens._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseCountTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseCountTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseCountTokens, + "_BaseCountTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1543,19 +1538,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectPredict._get_http_options() ) - request, metadata = self._interceptor.pre_direct_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectPredict, + "_BaseDirectPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1700,21 +1692,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_http_options() ) - request, metadata = self._interceptor.pre_direct_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectRawPredict, + "_BaseDirectRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1858,19 +1847,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseEmbedContent._get_http_options() ) - request, metadata = self._interceptor.pre_embed_content(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseEmbedContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseEmbedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseEmbedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseEmbedContent, + "_BaseEmbedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2014,23 +2000,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseExplain._get_http_options() ) - request, metadata = self._interceptor.pre_explain(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseExplain._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BaseExplain._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BaseExplain._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseExplain, + "_BaseExplain__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2174,21 +2153,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGenerateContent._get_http_options() ) - request, metadata = self._interceptor.pre_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGenerateContent, + "_BaseGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2332,23 +2308,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BasePredict._get_http_options() ) - request, metadata = self._interceptor.pre_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BasePredict._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BasePredict._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BasePredict._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BasePredict, + "_BasePredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2539,19 +2508,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseRawPredict._get_http_options() ) - request, metadata = self._interceptor.pre_raw_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseRawPredict, + "_BaseRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2700,21 +2666,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_http_options() ) - request, metadata = self._interceptor.pre_server_streaming_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseServerStreamingPredict, + "_BaseServerStreamingPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2891,21 +2854,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_http_options() ) - request, metadata = self._interceptor.pre_stream_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamGenerateContent, + "_BaseStreamGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3130,21 +3090,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_http_options() ) - request, metadata = self._interceptor.pre_stream_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamRawPredict, + "_BaseStreamRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3441,15 +3398,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3583,15 +3541,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3725,19 +3684,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3872,19 +3828,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4020,21 +3973,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4166,17 +4116,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4282,17 +4233,18 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4400,15 +4352,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4542,15 +4495,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4684,15 +4638,16 @@ def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_asyncio.py index a5138d322b..05a7d8b7f2 100644 --- a/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_asyncio.py @@ -36,6 +36,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1371,21 +1373,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseChatCompletions._get_http_options() ) - request, metadata = await self._interceptor.pre_chat_completions( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseChatCompletions._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseChatCompletions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseChatCompletions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseChatCompletions, + "_BaseChatCompletions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1528,21 +1527,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseCountTokens._get_http_options() ) - request, metadata = await self._interceptor.pre_count_tokens( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseCountTokens._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseCountTokens._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseCountTokens._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseCountTokens, + "_BaseCountTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1695,21 +1691,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_direct_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectPredict, + "_BaseDirectPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1862,21 +1855,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_direct_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDirectRawPredict, + "_BaseDirectRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2027,21 +2017,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseEmbedContent._get_http_options() ) - request, metadata = await self._interceptor.pre_embed_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseEmbedContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseEmbedContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseEmbedContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseEmbedContent, + "_BaseEmbedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2193,23 +2180,16 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseExplain._get_http_options() ) - request, metadata = await self._interceptor.pre_explain(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BaseExplain._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BaseExplain._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BaseExplain._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseExplain, + "_BaseExplain__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2359,21 +2339,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGenerateContent._get_http_options() ) - request, metadata = await self._interceptor.pre_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGenerateContent, + "_BaseGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2523,23 +2500,16 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BasePredict._get_http_options() ) - request, metadata = await self._interceptor.pre_predict(request, metadata) - transcoded_request = _BasePredictionServiceRestTransport._BasePredict._get_transcoded_request( - http_options, request - ) - - body = ( - _BasePredictionServiceRestTransport._BasePredict._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BasePredictionServiceRestTransport._BasePredict._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BasePredict, + "_BasePredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2737,21 +2707,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseRawPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseRawPredict, + "_BaseRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2907,21 +2874,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_server_streaming_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseServerStreamingPredict, + "_BaseServerStreamingPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3103,21 +3067,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_http_options() ) - request, metadata = await self._interceptor.pre_stream_generate_content( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamGenerateContent, + "_BaseStreamGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3347,21 +3308,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_http_options() ) - request, metadata = await self._interceptor.pre_stream_raw_predict( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseStreamRawPredict, + "_BaseStreamRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3635,17 +3593,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3787,17 +3746,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3940,21 +3900,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4098,21 +4055,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4256,21 +4210,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4408,17 +4359,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4530,17 +4482,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4655,17 +4608,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4807,17 +4761,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4959,17 +4914,18 @@ async def __call__( http_options = ( _BasePredictionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BasePredictionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BasePredictionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BasePredictionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_base.py index c9b8b9e9bb..521a94b0b9 100644 --- a/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/prediction_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.ChatCompletionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseChatCompletions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCountTokens: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -191,52 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.CountTokensRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseCountTokens._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDirectPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -248,52 +160,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.DirectPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseDirectPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDirectRawPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -305,38 +177,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.DirectRawPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseDirectRawPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseEmbedContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -352,47 +192,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.EmbedContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExplain: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -404,52 +209,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.ExplainRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseExplain._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGenerateContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -476,52 +241,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.GenerateContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseGenerateContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -538,52 +263,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.PredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BasePredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRawPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -600,52 +285,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.RawPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseRawPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseServerStreamingPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -662,38 +307,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.StreamingPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseServerStreamingPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamDirectPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -708,14 +321,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -742,38 +347,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.GenerateContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseStreamGenerateContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamingPredict: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -788,14 +361,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -812,38 +377,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = prediction_service.StreamRawPredictRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BasePredictionServiceRestTransport._BaseStreamRawPredict._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -862,17 +395,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -891,17 +413,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -985,22 +496,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1091,22 +586,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1182,22 +661,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1608,17 +1071,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2073,17 +1525,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2546,17 +1987,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3015,17 +2445,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3480,16 +2899,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BasePredictionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/client.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/client.py index a77c906548..d344d5a760 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/client.py @@ -40,6 +40,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -144,77 +151,13 @@ class ReasoningEngineExecutionServiceClient( ): """A service for executing queries on Reasoning Engine.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -402,9 +345,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - ReasoningEngineExecutionServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -425,37 +366,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - ReasoningEngineExecutionServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -475,67 +391,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -669,17 +524,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ReasoningEngineExecutionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( ReasoningEngineExecutionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - ReasoningEngineExecutionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -717,14 +572,14 @@ def __init__( self._transport = cast(ReasoningEngineExecutionServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ReasoningEngineExecutionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest.py index 634319ca3e..b3bc92176a 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2961,21 +2962,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_async_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine, + "_BaseAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3119,23 +3117,20 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = ( self._interceptor.pre_cancel_async_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine, + "_BaseCancelAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3288,21 +3283,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine, + "_BaseQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3500,21 +3492,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_stream_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine, + "_BaseStreamQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3694,15 +3683,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3838,15 +3828,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3981,19 +3972,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4129,19 +4117,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4279,21 +4264,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4425,17 +4407,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4541,17 +4524,18 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4660,15 +4644,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4802,15 +4787,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4944,15 +4930,16 @@ def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_asyncio.py index 98b6454451..7e04614856 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -811,23 +813,20 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = ( await self._interceptor.pre_async_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine, + "_BaseAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -981,23 +980,20 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_http_options() ) - request, metadata = ( await self._interceptor.pre_cancel_async_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine, + "_BaseCancelAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1158,21 +1154,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_query_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine, + "_BaseQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1375,23 +1368,20 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_http_options() ) - request, metadata = ( await self._interceptor.pre_stream_query_reasoning_engine( request, metadata ) ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine, + "_BaseStreamQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3796,17 +3786,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3948,17 +3939,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4101,21 +4093,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4259,21 +4248,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4417,21 +4403,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4571,17 +4554,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4695,17 +4679,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4822,17 +4807,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4974,17 +4960,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5126,17 +5113,18 @@ async def __call__( http_options = ( _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineExecutionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_base.py index 565b6cca09..074b83c1dc 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_execution_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -126,56 +118,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - reasoning_engine_execution_service.AsyncQueryReasoningEngineRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseAsyncQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCancelAsyncQueryReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -192,54 +140,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_execution_service.CancelAsyncQueryReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseCancelAsyncQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseQueryReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -261,56 +167,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - reasoning_engine_execution_service.QueryReasoningEngineRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStreamQueryReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -332,42 +194,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - reasoning_engine_execution_service.StreamQueryReasoningEngineRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineExecutionServiceRestTransport._BaseStreamQueryReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -386,17 +212,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -415,17 +230,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -509,22 +313,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -615,22 +403,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -706,22 +478,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1132,17 +888,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1597,17 +1342,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2070,17 +1804,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2539,17 +2262,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3004,16 +2716,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseReasoningEngineExecutionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/client.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/client.py index c2a7e59e4e..028031d761 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -151,77 +158,13 @@ class ReasoningEngineRuntimeRevisionServiceClient( ): """Manages Vertex AI's Reasoning Engine Revisions.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -479,9 +422,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ( - ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() - ) + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -502,37 +443,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ( - ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() - ) - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -552,69 +468,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ( - ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_UNIVERSE - ) - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ( - ReasoningEngineRuntimeRevisionServiceClient.DEFAULT_MTLS_ENDPOINT - ) - else: - api_endpoint = ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -748,17 +601,17 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ( ReasoningEngineRuntimeRevisionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) ) - self._universe_domain = ( - ReasoningEngineRuntimeRevisionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env - ) + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -798,14 +651,14 @@ def __init__( ) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ReasoningEngineRuntimeRevisionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest.py index 12ba5ad800..4432588c7a 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -2907,19 +2908,20 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision._get_http_options() ) - request, metadata = ( self._interceptor.pre_delete_reasoning_engine_runtime_revision( request, metadata ) ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision, + "_BaseDeleteReasoningEngineRuntimeRevision__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3065,19 +3067,20 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision._get_http_options() ) - request, metadata = ( self._interceptor.pre_get_reasoning_engine_runtime_revision( request, metadata ) ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision, + "_BaseGetReasoningEngineRuntimeRevision__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3231,19 +3234,20 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions._get_http_options() ) - request, metadata = ( self._interceptor.pre_list_reasoning_engine_runtime_revisions( request, metadata ) ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions, + "_BaseListReasoningEngineRuntimeRevisions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3433,15 +3437,16 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3577,15 +3582,16 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3722,19 +3728,16 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3872,19 +3875,16 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4022,21 +4022,18 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4170,17 +4167,18 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4288,17 +4286,18 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4409,15 +4408,16 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4553,15 +4553,16 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4697,15 +4698,16 @@ def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_asyncio.py index 86c4b308fa..4277dc6a44 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -749,19 +751,20 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_reasoning_engine_runtime_revision( request, metadata ) ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision, + "_BaseDeleteReasoningEngineRuntimeRevision__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -919,19 +922,20 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision._get_http_options() ) - request, metadata = ( await self._interceptor.pre_get_reasoning_engine_runtime_revision( request, metadata ) ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision, + "_BaseGetReasoningEngineRuntimeRevision__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1093,19 +1097,20 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_reasoning_engine_runtime_revisions( request, metadata ) ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions, + "_BaseListReasoningEngineRuntimeRevisions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3522,17 +3527,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3674,17 +3680,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3827,21 +3834,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3985,21 +3989,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4143,21 +4144,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4297,17 +4295,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4421,17 +4420,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4548,17 +4548,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4700,17 +4701,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4852,17 +4854,18 @@ async def __call__( http_options = ( _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_base.py index 36c038333b..40b334a6bf 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_runtime_revision_service/transports/rest_base.py @@ -104,14 +104,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -122,45 +114,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_runtime_revision_service.DeleteReasoningEngineRuntimeRevisionRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseDeleteReasoningEngineRuntimeRevision._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetReasoningEngineRuntimeRevision: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -171,45 +130,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_runtime_revision_service.GetReasoningEngineRuntimeRevisionRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseGetReasoningEngineRuntimeRevision._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListReasoningEngineRuntimeRevisions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -220,31 +146,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_runtime_revision_service.ListReasoningEngineRuntimeRevisionsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineRuntimeRevisionServiceRestTransport._BaseListReasoningEngineRuntimeRevisions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -263,17 +164,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -292,17 +182,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -386,22 +265,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -492,22 +355,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -583,22 +430,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1009,17 +840,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1474,17 +1294,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1947,17 +1756,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2416,17 +2214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2881,16 +2668,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseReasoningEngineRuntimeRevisionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/client.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/client.py index 2c18498dc0..70359edd1a 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -139,77 +146,13 @@ def get_transport_class( class ReasoningEngineServiceClient(metaclass=ReasoningEngineServiceClientMeta): """A service for managing Vertex AI's Reasoning Engines.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -489,7 +432,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ReasoningEngineServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -510,35 +453,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ReasoningEngineServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -558,67 +478,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ReasoningEngineServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ReasoningEngineServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -752,13 +611,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ReasoningEngineServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ReasoningEngineServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ReasoningEngineServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ReasoningEngineServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -794,14 +655,14 @@ def __init__( self._transport = cast(ReasoningEngineServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ReasoningEngineServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ReasoningEngineServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest.py index 4487c92e95..7df4222419 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3007,21 +3008,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_create_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine, + "_BaseCreateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3162,17 +3160,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_delete_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine, + "_BaseDeleteReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3313,17 +3312,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_get_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine, + "_BaseGetReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3468,17 +3468,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_http_options() ) - request, metadata = self._interceptor.pre_list_reasoning_engines( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines, + "_BaseListReasoningEngines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3627,21 +3628,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_http_options() ) - request, metadata = self._interceptor.pre_update_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine, + "_BaseUpdateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3836,15 +3834,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3978,15 +3977,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4121,19 +4121,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4269,19 +4266,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4417,21 +4411,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4565,17 +4556,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4683,17 +4675,18 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4804,15 +4797,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4946,15 +4940,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5090,15 +5085,16 @@ def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_asyncio.py index 88747e08ba..f41eec68e5 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -863,21 +865,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_create_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine, + "_BaseCreateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1030,17 +1029,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine, + "_BaseDeleteReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1191,17 +1191,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_get_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine, + "_BaseGetReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1350,17 +1351,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_http_options() ) - request, metadata = await self._interceptor.pre_list_reasoning_engines( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines, + "_BaseListReasoningEngines__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1515,21 +1517,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_http_options() ) - request, metadata = await self._interceptor.pre_update_reasoning_engine( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine, + "_BaseUpdateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3949,17 +3948,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4099,17 +4099,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4250,21 +4251,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4406,21 +4404,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4562,21 +4557,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4714,17 +4706,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4836,17 +4829,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4961,17 +4955,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5111,17 +5106,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5261,17 +5257,18 @@ async def __call__( http_options = ( _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseReasoningEngineServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseReasoningEngineServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_base.py index 9688c76ffc..ffa80d31d7 100644 --- a/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/reasoning_engine_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -124,54 +116,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.CreateReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseCreateReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -186,45 +136,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.DeleteReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseDeleteReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -239,43 +156,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.GetReasoningEngineRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseGetReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListReasoningEngines: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -290,45 +176,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.ListReasoningEnginesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseListReasoningEngines._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateReasoningEngine: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -345,40 +198,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = reasoning_engine_service.UpdateReasoningEngineRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseReasoningEngineServiceRestTransport._BaseUpdateReasoningEngine._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -397,17 +216,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -426,17 +234,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -520,22 +317,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -626,22 +407,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -717,22 +482,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1143,17 +892,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1608,17 +1346,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2081,17 +1808,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2550,17 +2266,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3015,16 +2720,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseReasoningEngineServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/schedule_service/client.py b/google/cloud/aiplatform_v1beta1/services/schedule_service/client.py index 13e35e118d..6e83a9dbf5 100644 --- a/google/cloud/aiplatform_v1beta1/services/schedule_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/schedule_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -142,77 +149,13 @@ class ScheduleServiceClient(metaclass=ScheduleServiceClientMeta): calls. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -757,7 +700,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = ScheduleServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -778,35 +721,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = ScheduleServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -826,65 +746,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = ScheduleServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = ScheduleServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -1016,13 +877,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - ScheduleServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = ScheduleServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = ScheduleServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=ScheduleServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -1058,14 +921,14 @@ def __init__( self._transport = cast(ScheduleServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or ScheduleServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=ScheduleServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest.py index a441ef0b82..7c24ee8366 100644 --- a/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3037,19 +3038,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_create_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCreateSchedule, + "_BaseCreateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3191,15 +3189,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_delete_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteSchedule, + "_BaseDeleteSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3339,15 +3338,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_get_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetSchedule, + "_BaseGetSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3487,15 +3487,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListSchedules._get_http_options() ) - request, metadata = self._interceptor.pre_list_schedules(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListSchedules._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListSchedules._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListSchedules, + "_BaseListSchedules__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3632,19 +3633,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BasePauseSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_pause_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BasePauseSchedule, + "_BasePauseSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3746,19 +3744,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_resume_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseResumeSchedule, + "_BaseResumeSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3868,19 +3863,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_http_options() ) - request, metadata = self._interceptor.pre_update_schedule(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseUpdateSchedule, + "_BaseUpdateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4079,15 +4071,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4220,15 +4213,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4362,19 +4356,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4509,19 +4500,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4657,21 +4645,18 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4802,17 +4787,18 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4917,17 +4903,18 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5035,15 +5022,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5176,15 +5164,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5317,15 +5306,16 @@ def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_asyncio.py index dc81d34dab..1e8e02f4a4 100644 --- a/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -900,21 +902,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_create_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCreateSchedule, + "_BaseCreateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1065,17 +1064,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteSchedule, + "_BaseDeleteSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1225,17 +1225,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_get_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetSchedule, + "_BaseGetSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1384,17 +1385,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListSchedules._get_http_options() ) - request, metadata = await self._interceptor.pre_list_schedules( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListSchedules._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListSchedules._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListSchedules, + "_BaseListSchedules__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1540,21 +1542,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BasePauseSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_pause_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BasePauseSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BasePauseSchedule, + "_BasePauseSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1665,21 +1664,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_resume_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseResumeSchedule, + "_BaseResumeSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1798,21 +1794,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_http_options() ) - request, metadata = await self._interceptor.pre_update_schedule( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseUpdateSchedule, + "_BaseUpdateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4230,17 +4223,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4382,17 +4376,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4535,21 +4530,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4693,21 +4685,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4851,21 +4840,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5003,17 +4989,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5127,17 +5114,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5254,17 +5242,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5406,17 +5395,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5558,17 +5548,18 @@ async def __call__( http_options = ( _BaseScheduleServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseScheduleServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseScheduleServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_base.py index ef0a624914..935313000c 100644 --- a/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/schedule_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.CreateScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseCreateSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -177,43 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.DeleteScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseDeleteSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -224,43 +145,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.GetScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseGetSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSchedules: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -271,43 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.ListSchedulesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseListSchedules._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BasePauseSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -319,52 +178,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.PauseScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BasePauseSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseResumeSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -376,38 +195,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.ResumeScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseResumeSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateSchedule: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -416,14 +203,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -435,38 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = schedule_service.UpdateScheduleRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseScheduleServiceRestTransport._BaseUpdateSchedule._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -485,17 +232,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -514,17 +250,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -608,22 +333,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -714,22 +423,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -805,22 +498,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1231,17 +908,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1696,17 +1362,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2169,17 +1824,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2638,17 +2282,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3103,16 +2736,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseScheduleServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/session_service/client.py b/google/cloud/aiplatform_v1beta1/services/session_service/client.py index 36d52a0af2..484e6843a3 100644 --- a/google/cloud/aiplatform_v1beta1/services/session_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/session_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -138,77 +145,13 @@ def get_transport_class( class SessionServiceClient(metaclass=SessionServiceClientMeta): """The service that manages Vertex Session related resources.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -446,7 +389,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = SessionServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -467,35 +410,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = SessionServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -515,65 +435,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = SessionServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = SessionServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = SessionServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -703,13 +564,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - SessionServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = SessionServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = SessionServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=SessionServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -745,14 +608,14 @@ def __init__( self._transport = cast(SessionServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or SessionServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=SessionServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=SessionServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest.py index 9283b5ffed..08b943bda7 100644 --- a/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3110,19 +3111,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseAppendEvent._get_http_options() ) - request, metadata = self._interceptor.pre_append_event(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseAppendEvent._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseAppendEvent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseAppendEvent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseAppendEvent, + "_BaseAppendEvent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3267,19 +3265,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCreateSession._get_http_options() ) - request, metadata = self._interceptor.pre_create_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseCreateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseCreateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCreateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCreateSession, + "_BaseCreateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3419,15 +3414,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteSession._get_http_options() ) - request, metadata = self._interceptor.pre_delete_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteSession, + "_BaseDeleteSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3565,17 +3561,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetSession._get_http_options() ) - request, metadata = self._interceptor.pre_get_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseGetSession._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetSession, + "_BaseGetSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3715,17 +3710,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListEvents._get_http_options() ) - request, metadata = self._interceptor.pre_list_events(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListEvents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseListEvents._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListEvents, + "_BaseListEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3867,15 +3861,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListSessions._get_http_options() ) - request, metadata = self._interceptor.pre_list_sessions(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListSessions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListSessions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListSessions, + "_BaseListSessions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4018,19 +4013,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseUpdateSession._get_http_options() ) - request, metadata = self._interceptor.pre_update_session(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseUpdateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseUpdateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseUpdateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseUpdateSession, + "_BaseUpdateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4233,15 +4225,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4374,15 +4367,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4516,19 +4510,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4663,19 +4654,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4810,21 +4798,18 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4955,17 +4940,18 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5070,17 +5056,18 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5188,15 +5175,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5329,15 +5317,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5470,15 +5459,16 @@ def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseSessionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_asyncio.py index fa875df357..dd29f7dd8c 100644 --- a/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -974,21 +976,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseAppendEvent._get_http_options() ) - request, metadata = await self._interceptor.pre_append_event( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseAppendEvent._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseAppendEvent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseAppendEvent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseAppendEvent, + "_BaseAppendEvent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1141,21 +1140,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCreateSession._get_http_options() ) - request, metadata = await self._interceptor.pre_create_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseCreateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseCreateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCreateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCreateSession, + "_BaseCreateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1305,17 +1301,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteSession._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteSession, + "_BaseDeleteSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1463,19 +1460,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetSession._get_http_options() ) - request, metadata = await self._interceptor.pre_get_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetSession._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseGetSession._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetSession, + "_BaseGetSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1621,19 +1617,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListEvents._get_http_options() ) - request, metadata = await self._interceptor.pre_list_events( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListEvents._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseSessionServiceRestTransport._BaseListEvents._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListEvents, + "_BaseListEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1781,17 +1776,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListSessions._get_http_options() ) - request, metadata = await self._interceptor.pre_list_sessions( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListSessions._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListSessions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListSessions, + "_BaseListSessions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1942,21 +1938,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseUpdateSession._get_http_options() ) - request, metadata = await self._interceptor.pre_update_session( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseUpdateSession._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseUpdateSession._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseUpdateSession._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseUpdateSession, + "_BaseUpdateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4378,17 +4371,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4529,17 +4523,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4681,21 +4676,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4838,21 +4830,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4996,21 +4985,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5148,17 +5134,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5272,17 +5259,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5398,17 +5386,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5550,17 +5539,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5701,17 +5691,18 @@ async def __call__( http_options = ( _BaseSessionServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseSessionServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSessionServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSessionServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_base.py index a3342a1936..8078bbb4c8 100644 --- a/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/session_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -125,52 +117,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.AppendEventRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseAppendEvent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -187,52 +139,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.CreateSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseCreateSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -247,43 +159,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.DeleteSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseDeleteSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -298,43 +179,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.GetSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseGetSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListEvents: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -349,43 +199,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.ListEventsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseListEvents._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSessions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -400,43 +219,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.ListSessionsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseListSessions._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateSession: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -453,38 +241,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = session_service.UpdateSessionRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSessionServiceRestTransport._BaseUpdateSession._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -503,17 +259,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -532,17 +277,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -626,22 +360,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -732,22 +450,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -823,22 +525,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1249,17 +935,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1714,17 +1389,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2187,17 +1851,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2656,17 +2309,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3121,16 +2763,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseSessionServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/client.py b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/client.py index 536684ae99..9230f3e42f 100644 --- a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -141,77 +148,13 @@ class SpecialistPoolServiceClient(metaclass=SpecialistPoolServiceClientMeta): manage Specialists and tasks on CrowdCompute console. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -399,7 +342,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = SpecialistPoolServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -420,35 +363,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = SpecialistPoolServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -468,67 +388,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = SpecialistPoolServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = ( - SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = SpecialistPoolServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -662,13 +521,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - SpecialistPoolServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = SpecialistPoolServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = SpecialistPoolServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=SpecialistPoolServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -704,14 +565,14 @@ def __init__( self._transport = cast(SpecialistPoolServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or SpecialistPoolServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=SpecialistPoolServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest.py index 96c1162b6b..ac308d2a4b 100644 --- a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3010,21 +3011,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_create_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool, + "_BaseCreateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3167,17 +3165,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_delete_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool, + "_BaseDeleteSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3328,17 +3327,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_get_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool, + "_BaseGetSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3481,17 +3481,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_http_options() ) - request, metadata = self._interceptor.pre_list_specialist_pools( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools, + "_BaseListSpecialistPools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3640,21 +3641,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_http_options() ) - request, metadata = self._interceptor.pre_update_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool, + "_BaseUpdateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3848,15 +3846,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3990,15 +3989,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4133,19 +4133,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4281,19 +4278,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4429,21 +4423,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4577,17 +4568,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4695,17 +4687,18 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4816,15 +4809,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4958,15 +4952,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5100,15 +5095,16 @@ def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_asyncio.py index d52cdc0b52..1a4669baa7 100644 --- a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -862,21 +864,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_create_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool, + "_BaseCreateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1025,17 +1024,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool, + "_BaseDeleteSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1192,17 +1192,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_get_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool, + "_BaseGetSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1349,17 +1350,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_http_options() ) - request, metadata = await self._interceptor.pre_list_specialist_pools( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools, + "_BaseListSpecialistPools__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1512,21 +1514,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_http_options() ) - request, metadata = await self._interceptor.pre_update_specialist_pool( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool, + "_BaseUpdateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3941,17 +3940,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4091,17 +4091,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4242,21 +4243,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4398,21 +4396,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4554,21 +4549,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4706,17 +4698,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4828,17 +4821,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4953,17 +4947,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5103,17 +5098,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5253,17 +5249,18 @@ async def __call__( http_options = ( _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseSpecialistPoolServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_base.py index 4c955e5a5c..fe3104083b 100644 --- a/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/specialist_pool_service/transports/rest_base.py @@ -100,14 +100,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -119,52 +111,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.CreateSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseCreateSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteSpecialistPool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,43 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.DeleteSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseDeleteSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetSpecialistPool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -222,43 +143,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.GetSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseGetSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListSpecialistPools: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -269,29 +159,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.ListSpecialistPoolsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseListSpecialistPools._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateSpecialistPool: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -300,14 +167,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -319,38 +178,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = specialist_pool_service.UpdateSpecialistPoolRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseSpecialistPoolServiceRestTransport._BaseUpdateSpecialistPool._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -369,17 +196,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -398,17 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -492,22 +297,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -598,22 +387,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -689,22 +462,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1115,17 +872,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1580,17 +1326,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2053,17 +1788,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2522,17 +2246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2987,16 +2700,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseSpecialistPoolServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/client.py b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/client.py index a96f936e48..10531cc564 100644 --- a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/client.py @@ -40,6 +40,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -149,77 +156,13 @@ def get_transport_class( class TensorboardServiceClient(metaclass=TensorboardServiceClientMeta): """TensorboardService""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -487,7 +430,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = TensorboardServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -508,35 +451,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = TensorboardServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -556,65 +476,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = TensorboardServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = TensorboardServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -748,13 +609,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - TensorboardServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = TensorboardServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = TensorboardServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=TensorboardServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -790,14 +653,14 @@ def __init__( self._transport = cast(TensorboardServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or TensorboardServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=TensorboardServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest.py index e8090fe94a..b4b6627dd7 100644 --- a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -4500,21 +4501,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_tensorboard_runs( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns, + "_BaseBatchCreateTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4666,23 +4664,20 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_create_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries, + "_BaseBatchCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4833,19 +4828,20 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( self._interceptor.pre_batch_read_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData, + "_BaseBatchReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4994,21 +4990,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboard, + "_BaseCreateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5152,21 +5145,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment, + "_BaseCreateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5317,21 +5307,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun, + "_BaseCreateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5479,21 +5466,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_create_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries, + "_BaseCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5642,17 +5626,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard, + "_BaseDeleteTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5793,17 +5778,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment, + "_BaseDeleteTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5945,17 +5931,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun, + "_BaseDeleteTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6098,17 +6085,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_delete_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries, + "_BaseDeleteTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6253,23 +6241,20 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( self._interceptor.pre_export_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData, + "_BaseExportTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6421,15 +6406,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboard, + "_BaseGetTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6573,17 +6559,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment, + "_BaseGetTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6730,17 +6717,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun, + "_BaseGetTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6882,17 +6870,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_get_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries, + "_BaseGetTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7038,17 +7027,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboard_experiments( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments, + "_BaseListTensorboardExperiments__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7193,17 +7183,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboard_runs( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns, + "_BaseListTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7350,17 +7341,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboards( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboards, + "_BaseListTensorboards__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7504,17 +7496,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_list_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries, + "_BaseListTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7661,17 +7654,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_blob_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData, + "_BaseReadTensorboardBlobData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7808,17 +7802,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_size( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize, + "_BaseReadTensorboardSize__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7966,17 +7961,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_time_series_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData, + "_BaseReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8121,17 +8117,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_http_options() ) - request, metadata = self._interceptor.pre_read_tensorboard_usage( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage, + "_BaseReadTensorboardUsage__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8280,21 +8277,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard, + "_BaseUpdateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8438,21 +8432,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment, + "_BaseUpdateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8603,21 +8594,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun, + "_BaseUpdateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8765,21 +8753,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_http_options() ) - request, metadata = self._interceptor.pre_update_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries, + "_BaseUpdateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8931,21 +8916,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_http_options() ) - request, metadata = self._interceptor.pre_write_tensorboard_experiment_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData, + "_BaseWriteTensorboardExperimentData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9095,21 +9077,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_http_options() ) - request, metadata = self._interceptor.pre_write_tensorboard_run_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData, + "_BaseWriteTensorboardRunData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9579,15 +9558,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9721,15 +9701,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9864,19 +9845,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10012,19 +9990,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10160,21 +10135,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10308,17 +10280,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10424,17 +10397,18 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10543,15 +10517,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10685,15 +10660,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10827,15 +10803,16 @@ def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_asyncio.py index 62cee79d37..970649db40 100644 --- a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -2480,23 +2482,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_create_tensorboard_runs( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns, + "_BaseBatchCreateTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2654,23 +2653,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_create_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries, + "_BaseBatchCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2829,19 +2825,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_read_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData, + "_BaseBatchReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2998,21 +2995,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_create_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboard, + "_BaseCreateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3166,23 +3160,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_tensorboard_experiment( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment, + "_BaseCreateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3339,21 +3330,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_create_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun, + "_BaseCreateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3507,23 +3495,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_create_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries, + "_BaseCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3678,17 +3663,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard, + "_BaseDeleteTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3839,19 +3825,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_tensorboard_experiment( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment, + "_BaseDeleteTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4001,17 +3988,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun, + "_BaseDeleteTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4162,19 +4150,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_delete_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries, + "_BaseDeleteTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4327,23 +4316,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_export_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData, + "_BaseExportTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4503,17 +4489,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboard, + "_BaseGetTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4663,17 +4650,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard_experiment( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment, + "_BaseGetTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4826,17 +4814,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun, + "_BaseGetTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4984,17 +4973,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_http_options() ) - request, metadata = await self._interceptor.pre_get_tensorboard_time_series( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries, + "_BaseGetTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5148,19 +5138,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_tensorboard_experiments( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments, + "_BaseListTensorboardExperiments__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5313,17 +5304,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_http_options() ) - request, metadata = await self._interceptor.pre_list_tensorboard_runs( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns, + "_BaseListTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5474,17 +5466,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_http_options() ) - request, metadata = await self._interceptor.pre_list_tensorboards( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboards, + "_BaseListTensorboards__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5636,19 +5629,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_list_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries, + "_BaseListTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5802,17 +5796,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_http_options() ) - request, metadata = await self._interceptor.pre_read_tensorboard_blob_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData, + "_BaseReadTensorboardBlobData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5955,17 +5950,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_http_options() ) - request, metadata = await self._interceptor.pre_read_tensorboard_size( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize, + "_BaseReadTensorboardSize__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6119,19 +6115,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_read_tensorboard_time_series_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData, + "_BaseReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6282,17 +6279,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_http_options() ) - request, metadata = await self._interceptor.pre_read_tensorboard_usage( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage, + "_BaseReadTensorboardUsage__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6445,21 +6443,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_http_options() ) - request, metadata = await self._interceptor.pre_update_tensorboard( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard, + "_BaseUpdateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6613,23 +6608,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_tensorboard_experiment( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment, + "_BaseUpdateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6786,21 +6778,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_http_options() ) - request, metadata = await self._interceptor.pre_update_tensorboard_run( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun, + "_BaseUpdateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6954,23 +6943,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_http_options() ) - request, metadata = ( await self._interceptor.pre_update_tensorboard_time_series( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries, + "_BaseUpdateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7128,23 +7114,20 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_http_options() ) - request, metadata = ( await self._interceptor.pre_write_tensorboard_experiment_data( request, metadata ) ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData, + "_BaseWriteTensorboardExperimentData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7300,21 +7283,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_http_options() ) - request, metadata = await self._interceptor.pre_write_tensorboard_run_data( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData, + "_BaseWriteTensorboardRunData__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9955,17 +9935,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10107,17 +10088,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10260,21 +10242,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10418,21 +10397,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10576,21 +10552,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10728,17 +10701,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10850,17 +10824,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -10975,17 +10950,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11127,17 +11103,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -11277,17 +11254,18 @@ async def __call__( http_options = ( _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseTensorboardServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseTensorboardServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_base.py index f29ca76f31..15ba52435b 100644 --- a/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/tensorboard_service/transports/rest_base.py @@ -110,14 +110,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -129,54 +121,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.BatchCreateTensorboardRunsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardRuns._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchCreateTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -188,40 +138,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.BatchCreateTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseBatchCreateTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchReadTensorboardTimeSeriesData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -230,14 +146,6 @@ def __hash__(self): # pragma: NO COVER "timeSeries": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -248,47 +156,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = ( - tensorboard_service.BatchReadTensorboardTimeSeriesDataRequest.pb( - request - ) - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseBatchReadTensorboardTimeSeriesData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -300,38 +173,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -340,14 +181,6 @@ def __hash__(self): # pragma: NO COVER "tensorboardExperimentId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -359,40 +192,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardExperimentRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -401,14 +200,6 @@ def __hash__(self): # pragma: NO COVER "tensorboardRunId": "", } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -420,52 +211,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -477,54 +228,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.CreateTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseCreateTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -535,43 +244,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -582,45 +260,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardExperimentRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -631,43 +276,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -678,45 +292,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.DeleteTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseDeleteTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseExportTensorboardTimeSeriesData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -728,54 +309,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ExportTensorboardTimeSeriesDataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseExportTensorboardTimeSeriesData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -786,43 +325,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -833,43 +341,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardExperimentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -880,43 +357,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -927,43 +373,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.GetTensorboardTimeSeriesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseGetTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboardExperiments: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -974,45 +389,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardExperimentsRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboardExperiments._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboardRuns: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1023,43 +405,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardRunsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboardRuns._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboards: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1070,43 +421,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboards._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1117,45 +437,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ListTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseListTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardBlobData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1166,43 +453,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardBlobDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardBlobData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardSize: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1213,43 +469,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardSizeRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardSize._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardTimeSeriesData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1260,45 +485,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardTimeSeriesDataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardTimeSeriesData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseReadTensorboardUsage: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1309,29 +501,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.ReadTensorboardUsageRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseReadTensorboardUsage._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboard: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1340,14 +509,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1359,38 +520,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboard._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboardExperiment: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1399,14 +528,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1418,40 +539,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardExperimentRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardExperiment._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboardRun: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1460,14 +547,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1479,38 +558,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardRunRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardRun._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateTensorboardTimeSeries: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1519,14 +566,6 @@ def __hash__(self): # pragma: NO COVER "updateMask": {}, } - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1538,54 +577,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.UpdateTensorboardTimeSeriesRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseUpdateTensorboardTimeSeries._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseWriteTensorboardExperimentData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1597,54 +594,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.WriteTensorboardExperimentDataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseWriteTensorboardExperimentData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseWriteTensorboardRunData: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1656,38 +611,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = tensorboard_service.WriteTensorboardRunDataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseTensorboardServiceRestTransport._BaseWriteTensorboardRunData._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1706,17 +629,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1735,17 +647,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1829,22 +730,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1935,22 +820,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2026,22 +895,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2452,17 +1305,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2917,17 +1759,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3390,17 +2221,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3859,17 +2679,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4324,16 +3133,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseTensorboardServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/client.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/client.py index 6d24768dfc..5af2c0df73 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -137,77 +144,13 @@ def get_transport_class( class VertexRagDataServiceClient(metaclass=VertexRagDataServiceClientMeta): """A service for managing user data for RAG.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -555,7 +498,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = VertexRagDataServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -576,35 +519,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = VertexRagDataServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -624,65 +544,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = VertexRagDataServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = VertexRagDataServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -816,13 +677,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - VertexRagDataServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = VertexRagDataServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = VertexRagDataServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=VertexRagDataServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -858,14 +721,14 @@ def __init__( self._transport = cast(VertexRagDataServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or VertexRagDataServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=VertexRagDataServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest.py index 1d8194d810..e7547efc08 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -4085,21 +4086,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_rag_data_schemas( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas, + "_BaseBatchCreateRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4243,21 +4241,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_http_options() ) - request, metadata = self._interceptor.pre_batch_create_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata, + "_BaseBatchCreateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4402,21 +4397,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_http_options() ) - request, metadata = self._interceptor.pre_batch_delete_rag_data_schemas( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas, + "_BaseBatchDeleteRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4560,21 +4552,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_http_options() ) - request, metadata = self._interceptor.pre_batch_delete_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata, + "_BaseBatchDeleteRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4718,21 +4707,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_create_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus, + "_BaseCreateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4873,21 +4859,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_http_options() ) - request, metadata = self._interceptor.pre_create_rag_data_schema( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema, + "_BaseCreateRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5032,21 +5015,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_http_options() ) - request, metadata = self._interceptor.pre_create_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata, + "_BaseCreateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5191,17 +5171,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_delete_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus, + "_BaseDeleteRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5334,17 +5315,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema._get_http_options() ) - request, metadata = self._interceptor.pre_delete_rag_data_schema( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema, + "_BaseDeleteRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5454,15 +5436,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_http_options() ) - request, metadata = self._interceptor.pre_delete_rag_file(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile, + "_BaseDeleteRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5595,17 +5578,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata._get_http_options() ) - request, metadata = self._interceptor.pre_delete_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata, + "_BaseDeleteRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5715,15 +5699,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_corpus(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus, + "_BaseGetRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5864,17 +5849,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_data_schema( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema, + "_BaseGetRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6015,17 +6001,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig, + "_BaseGetRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6168,15 +6155,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_file(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagFile, + "_BaseGetRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6317,17 +6305,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata._get_http_options() ) - request, metadata = self._interceptor.pre_get_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata, + "_BaseGetRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6470,21 +6459,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_http_options() ) - request, metadata = self._interceptor.pre_import_rag_files( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles, + "_BaseImportRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6624,17 +6610,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_http_options() ) - request, metadata = self._interceptor.pre_list_rag_corpora( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora, + "_BaseListRagCorpora__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6777,17 +6764,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas._get_http_options() ) - request, metadata = self._interceptor.pre_list_rag_data_schemas( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas, + "_BaseListRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6934,15 +6922,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_http_options() ) - request, metadata = self._interceptor.pre_list_rag_files(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagFiles, + "_BaseListRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7085,17 +7074,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata._get_http_options() ) - request, metadata = self._interceptor.pre_list_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata, + "_BaseListRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7242,21 +7232,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_http_options() ) - request, metadata = self._interceptor.pre_update_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus, + "_BaseUpdateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7398,21 +7385,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_http_options() ) - request, metadata = self._interceptor.pre_update_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig, + "_BaseUpdateRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7555,21 +7539,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_http_options() ) - request, metadata = self._interceptor.pre_update_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata, + "_BaseUpdateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7714,19 +7695,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_http_options() ) - request, metadata = self._interceptor.pre_upload_rag_file(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile, + "_BaseUploadRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8127,15 +8105,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8269,15 +8248,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8412,19 +8392,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8560,19 +8537,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8708,21 +8682,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8856,17 +8827,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8972,17 +8944,18 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9091,15 +9064,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9233,15 +9207,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9375,15 +9350,16 @@ def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_asyncio.py index 0830986169..093e1ee75d 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -2039,23 +2041,20 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_create_rag_data_schemas( request, metadata ) ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas, + "_BaseBatchCreateRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2207,21 +2206,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_create_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata, + "_BaseBatchCreateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2376,23 +2372,20 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_http_options() ) - request, metadata = ( await self._interceptor.pre_batch_delete_rag_data_schemas( request, metadata ) ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas, + "_BaseBatchDeleteRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2544,21 +2537,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_http_options() ) - request, metadata = await self._interceptor.pre_batch_delete_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata, + "_BaseBatchDeleteRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2710,21 +2700,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_create_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus, + "_BaseCreateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2873,21 +2860,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_http_options() ) - request, metadata = await self._interceptor.pre_create_rag_data_schema( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema, + "_BaseCreateRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3036,21 +3020,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_http_options() ) - request, metadata = await self._interceptor.pre_create_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata, + "_BaseCreateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3199,17 +3180,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus, + "_BaseDeleteRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3350,17 +3332,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_rag_data_schema( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema, + "_BaseDeleteRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3474,17 +3457,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_rag_file( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile, + "_BaseDeleteRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3625,17 +3609,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata, + "_BaseDeleteRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3749,17 +3734,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus, + "_BaseGetRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3906,17 +3892,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_data_schema( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema, + "_BaseGetRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4061,17 +4048,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig, + "_BaseGetRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4218,17 +4206,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_file( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagFile, + "_BaseGetRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4377,17 +4366,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata._get_http_options() ) - request, metadata = await self._interceptor.pre_get_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata, + "_BaseGetRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4536,21 +4526,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_http_options() ) - request, metadata = await self._interceptor.pre_import_rag_files( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles, + "_BaseImportRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4698,17 +4685,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_http_options() ) - request, metadata = await self._interceptor.pre_list_rag_corpora( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora, + "_BaseListRagCorpora__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4857,17 +4845,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas._get_http_options() ) - request, metadata = await self._interceptor.pre_list_rag_data_schemas( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas, + "_BaseListRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5018,17 +5007,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_http_options() ) - request, metadata = await self._interceptor.pre_list_rag_files( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagFiles, + "_BaseListRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5177,17 +5167,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata._get_http_options() ) - request, metadata = await self._interceptor.pre_list_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata, + "_BaseListRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5340,21 +5331,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_http_options() ) - request, metadata = await self._interceptor.pre_update_rag_corpus( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus, + "_BaseUpdateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5504,21 +5492,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_http_options() ) - request, metadata = await self._interceptor.pre_update_rag_engine_config( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig, + "_BaseUpdateRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5669,21 +5654,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_http_options() ) - request, metadata = await self._interceptor.pre_update_rag_metadata( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata, + "_BaseUpdateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5832,21 +5814,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_http_options() ) - request, metadata = await self._interceptor.pre_upload_rag_file( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile, + "_BaseUploadRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8428,17 +8407,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8580,17 +8560,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8731,21 +8712,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -8887,21 +8865,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9043,21 +9018,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9195,17 +9167,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9317,17 +9290,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9442,17 +9416,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9592,17 +9567,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -9742,17 +9718,18 @@ async def __call__( http_options = ( _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagDataServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagDataServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_base.py index 2355689296..acffa17566 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_data_service/transports/rest_base.py @@ -101,14 +101,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -120,54 +112,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.BatchCreateRagDataSchemasRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagDataSchemas._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchCreateRagMetadata: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -179,54 +129,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.BatchCreateRagMetadataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseBatchCreateRagMetadata._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchDeleteRagDataSchemas: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -238,54 +146,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.BatchDeleteRagDataSchemasRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagDataSchemas._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseBatchDeleteRagMetadata: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -297,54 +163,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.BatchDeleteRagMetadataRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseBatchDeleteRagMetadata._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateRagCorpus: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -356,52 +180,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.CreateRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseCreateRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateRagDataSchema: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -413,52 +197,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.CreateRagDataSchemaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseCreateRagDataSchema._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateRagMetadata: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -470,52 +214,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.CreateRagMetadataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseCreateRagMetadata._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteRagCorpus: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -526,43 +230,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.DeleteRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseDeleteRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteRagDataSchema: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -573,43 +246,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.DeleteRagDataSchemaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseDeleteRagDataSchema._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteRagFile: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -620,43 +262,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.DeleteRagFileRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseDeleteRagFile._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteRagMetadata: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -667,43 +278,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.DeleteRagMetadataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseDeleteRagMetadata._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagCorpus: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -714,43 +294,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagDataSchema: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -761,43 +310,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagDataSchemaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagDataSchema._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagEngineConfig: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -808,43 +326,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagEngineConfigRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagEngineConfig._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagFile: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -855,43 +342,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagFileRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagFile._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetRagMetadata: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -902,43 +358,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.GetRagMetadataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseGetRagMetadata._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseImportRagFiles: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -950,52 +375,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ImportRagFilesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseImportRagFiles._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListRagCorpora: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1006,43 +391,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ListRagCorporaRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseListRagCorpora._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListRagDataSchemas: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1053,43 +407,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ListRagDataSchemasRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseListRagDataSchemas._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListRagFiles: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1100,43 +423,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ListRagFilesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseListRagFiles._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListRagMetadata: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1147,43 +439,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.ListRagMetadataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseListRagMetadata._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateRagCorpus: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1195,52 +456,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.UpdateRagCorpusRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseUpdateRagCorpus._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateRagEngineConfig: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1252,54 +473,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.UpdateRagEngineConfigRequest.pb( - request - ) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseUpdateRagEngineConfig._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUpdateRagMetadata: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1311,52 +490,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.UpdateRagMetadataRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseUpdateRagMetadata._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseUploadRagFile: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -1368,38 +507,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_data_service.UploadRagFileRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagDataServiceRestTransport._BaseUploadRagFile._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1418,17 +525,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1447,17 +543,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1541,22 +626,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1647,22 +716,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1738,22 +791,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2164,17 +1201,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2629,17 +1655,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3102,17 +2117,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3571,17 +2575,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -4036,16 +3029,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseVertexRagDataServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/client.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/client.py index b1ccffabed..29b354a60e 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -132,77 +139,13 @@ def get_transport_class( class VertexRagServiceClient(metaclass=VertexRagServiceClientMeta): """A service for retrieving relevant contexts.""" - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -390,7 +333,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = VertexRagServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -411,35 +354,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = VertexRagServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -459,65 +379,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = VertexRagServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = VertexRagServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -649,13 +510,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - VertexRagServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = VertexRagServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = VertexRagServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=VertexRagServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -691,14 +554,14 @@ def __init__( self._transport = cast(VertexRagServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or VertexRagServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=VertexRagServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest.py index 14fb2aa445..948b9a6864 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3008,19 +3009,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAskContexts._get_http_options() ) - request, metadata = self._interceptor.pre_ask_contexts(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAskContexts, + "_BaseAskContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3166,21 +3164,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_http_options() ) - request, metadata = self._interceptor.pre_async_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts, + "_BaseAsyncRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3319,19 +3314,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_http_options() ) - request, metadata = self._interceptor.pre_augment_prompt(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAugmentPrompt, + "_BaseAugmentPrompt__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3476,21 +3468,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_http_options() ) - request, metadata = self._interceptor.pre_corroborate_content( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCorroborateContent, + "_BaseCorroborateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3635,21 +3624,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_http_options() ) - request, metadata = self._interceptor.pre_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseRetrieveContexts, + "_BaseRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3845,15 +3831,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3986,15 +3973,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4128,19 +4116,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4275,19 +4260,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4423,21 +4405,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4569,17 +4548,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4685,17 +4665,18 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4803,15 +4784,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4944,15 +4926,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5085,15 +5068,16 @@ def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_asyncio.py index f0c76cf76d..b18325e132 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -861,21 +863,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAskContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_ask_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAskContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAskContexts, + "_BaseAskContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1029,21 +1028,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_async_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts, + "_BaseAsyncRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1191,21 +1187,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_http_options() ) - request, metadata = await self._interceptor.pre_augment_prompt( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseAugmentPrompt, + "_BaseAugmentPrompt__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1358,21 +1351,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_http_options() ) - request, metadata = await self._interceptor.pre_corroborate_content( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCorroborateContent, + "_BaseCorroborateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1523,21 +1513,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_http_options() ) - request, metadata = await self._interceptor.pre_retrieve_contexts( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseRetrieveContexts, + "_BaseRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3955,17 +3942,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4107,17 +4095,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4260,21 +4249,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4418,21 +4404,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4576,21 +4559,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4728,17 +4708,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4852,17 +4833,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4979,17 +4961,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5131,17 +5114,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5283,17 +5267,18 @@ async def __call__( http_options = ( _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVertexRagServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVertexRagServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_base.py index 74fcc2131d..feef621223 100644 --- a/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/vertex_rag_service/transports/rest_base.py @@ -99,14 +99,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -118,52 +110,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.AskContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseAskContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAsyncRetrieveContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -175,52 +127,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.AsyncRetrieveContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseAsyncRetrieveContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseAugmentPrompt: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -232,52 +144,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.AugmentPromptRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseAugmentPrompt._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCorroborateContent: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -289,52 +161,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.CorroborateContentRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseCorroborateContent._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseRetrieveContexts: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -346,38 +178,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vertex_rag_service.RetrieveContextsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVertexRagServiceRestTransport._BaseRetrieveContexts._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -396,17 +196,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -425,17 +214,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -519,22 +297,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -625,22 +387,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -716,22 +462,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1142,17 +872,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1607,17 +1326,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2080,17 +1788,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2549,17 +2246,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3014,16 +2700,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseVertexRagServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/services/vizier_service/client.py b/google/cloud/aiplatform_v1beta1/services/vizier_service/client.py index d8ac77243f..7c8f8052d2 100644 --- a/google/cloud/aiplatform_v1beta1/services/vizier_service/client.py +++ b/google/cloud/aiplatform_v1beta1/services/vizier_service/client.py @@ -39,6 +39,13 @@ 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.aiplatform_v1beta1._compat import ( + get_universe_domain, + get_api_endpoint, + get_default_mtls_endpoint, + should_use_client_cert, + read_environment_variables, +) 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 @@ -136,77 +143,13 @@ class VizierServiceClient(metaclass=VizierServiceClientMeta): searching over deep learning architectures. """ - @staticmethod - def _get_default_mtls_endpoint(api_endpoint) -> Optional[str]: - """Converts api endpoint to mTLS endpoint. - - Convert "*.sandbox.googleapis.com" and "*.googleapis.com" to - "*.mtls.sandbox.googleapis.com" and "*.mtls.googleapis.com" respectively. - Args: - api_endpoint (Optional[str]): the api endpoint to convert. - Returns: - Optional[str]: converted mTLS api endpoint. - """ - if not api_endpoint: - return api_endpoint - - mtls_endpoint_re = re.compile( - r"(?P[^.]+)(?P\.mtls)?(?P\.sandbox)?(?P\.googleapis\.com)?" - ) - - m = mtls_endpoint_re.match(api_endpoint) - if m is None: - # Could not parse api_endpoint; return as-is. - return api_endpoint - - name, mtls, sandbox, googledomain = m.groups() - if mtls or not googledomain: - return api_endpoint - - if sandbox: - return api_endpoint.replace( - "sandbox.googleapis.com", "mtls.sandbox.googleapis.com" - ) - - return api_endpoint.replace(".googleapis.com", ".mtls.googleapis.com") - # Note: DEFAULT_ENDPOINT is deprecated. Use _DEFAULT_ENDPOINT_TEMPLATE instead. DEFAULT_ENDPOINT = "aiplatform.googleapis.com" - DEFAULT_MTLS_ENDPOINT = _get_default_mtls_endpoint.__func__( # type: ignore - DEFAULT_ENDPOINT - ) + DEFAULT_MTLS_ENDPOINT = get_default_mtls_endpoint(DEFAULT_ENDPOINT) _DEFAULT_ENDPOINT_TEMPLATE = "aiplatform.{UNIVERSE_DOMAIN}" _DEFAULT_UNIVERSE = "googleapis.com" - @staticmethod - def _use_client_cert_effective(): - """Returns whether client certificate should be used for mTLS if the - google-auth version supports should_use_client_cert automatic mTLS enablement. - - Alternatively, read from the GOOGLE_API_USE_CLIENT_CERTIFICATE env var. - - Returns: - bool: whether client certificate should be used for mTLS - Raises: - ValueError: (If using a version of google-auth without should_use_client_cert and - GOOGLE_API_USE_CLIENT_CERTIFICATE is set to an unexpected value.) - """ - # check if google-auth version supports should_use_client_cert for automatic mTLS enablement - if hasattr(mtls, "should_use_client_cert"): # pragma: NO COVER - return mtls.should_use_client_cert() - else: # pragma: NO COVER - # if unsupported, fallback to reading from env var - use_client_cert_str = os.getenv( - "GOOGLE_API_USE_CLIENT_CERTIFICATE", "false" - ).lower() - if use_client_cert_str not in ("true", "false"): - raise ValueError( - "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be" - " either `true` or `false`" - ) - return use_client_cert_str == "true" - @classmethod def from_service_account_info(cls, info: dict, *args, **kwargs): """Creates an instance of this client using the provided credentials @@ -440,7 +383,7 @@ def get_mtls_endpoint_and_cert_source( ) if client_options is None: client_options = client_options_lib.ClientOptions() - use_client_cert = VizierServiceClient._use_client_cert_effective() + use_client_cert = should_use_client_cert() use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto") if use_mtls_endpoint not in ("auto", "never", "always"): raise MutualTLSChannelError( @@ -461,35 +404,12 @@ def get_mtls_endpoint_and_cert_source( elif use_mtls_endpoint == "always" or ( use_mtls_endpoint == "auto" and client_cert_source ): - api_endpoint = cls.DEFAULT_MTLS_ENDPOINT + api_endpoint = cls.DEFAULT_MTLS_ENDPOINT # type: ignore else: api_endpoint = cls.DEFAULT_ENDPOINT return api_endpoint, client_cert_source - @staticmethod - def _read_environment_variables(): - """Returns the environment variables used by the client. - - Returns: - Tuple[bool, str, str]: returns the GOOGLE_API_USE_CLIENT_CERTIFICATE, - GOOGLE_API_USE_MTLS_ENDPOINT, and GOOGLE_CLOUD_UNIVERSE_DOMAIN environment variables. - - Raises: - ValueError: If GOOGLE_API_USE_CLIENT_CERTIFICATE is not - any of ["true", "false"]. - google.auth.exceptions.MutualTLSChannelError: If GOOGLE_API_USE_MTLS_ENDPOINT - is not any of ["auto", "never", "always"]. - """ - use_client_cert = VizierServiceClient._use_client_cert_effective() - use_mtls_endpoint = os.getenv("GOOGLE_API_USE_MTLS_ENDPOINT", "auto").lower() - universe_domain_env = os.getenv("GOOGLE_CLOUD_UNIVERSE_DOMAIN") - if use_mtls_endpoint not in ("auto", "never", "always"): - 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. @@ -509,65 +429,6 @@ def _get_client_cert_source(provided_cert_source, use_cert_flag): client_cert_source = mtls.default_client_cert_source() return client_cert_source - @staticmethod - def _get_api_endpoint( - api_override, client_cert_source, universe_domain, use_mtls_endpoint - ) -> str: - """Return the API endpoint used by the client. - - Args: - api_override (str): The API endpoint override. If specified, this is always - the return value of this function and the other arguments are not used. - client_cert_source (bytes): The client certificate source used by the client. - universe_domain (str): The universe domain used by the client. - use_mtls_endpoint (str): How to use the mTLS endpoint, which depends also on the other parameters. - Possible values are "always", "auto", or "never". - - Returns: - str: The API endpoint to be used by the client. - """ - if api_override is not None: - api_endpoint = api_override - elif use_mtls_endpoint == "always" or ( - use_mtls_endpoint == "auto" and client_cert_source - ): - _default_universe = VizierServiceClient._DEFAULT_UNIVERSE - if universe_domain != _default_universe: - raise MutualTLSChannelError( - f"mTLS is not supported in any universe other than {_default_universe}." - ) - api_endpoint = VizierServiceClient.DEFAULT_MTLS_ENDPOINT - else: - api_endpoint = VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format( - UNIVERSE_DOMAIN=universe_domain - ) - return api_endpoint - - @staticmethod - def _get_universe_domain( - client_universe_domain: Optional[str], universe_domain_env: Optional[str] - ) -> str: - """Return the universe domain used by the client. - - Args: - client_universe_domain (Optional[str]): The universe domain configured via the client options. - universe_domain_env (Optional[str]): The universe domain configured via the "GOOGLE_CLOUD_UNIVERSE_DOMAIN" environment variable. - - Returns: - str: The universe domain to be used by the client. - - Raises: - ValueError: If the universe domain is an empty string. - """ - universe_domain = VizierServiceClient._DEFAULT_UNIVERSE - if client_universe_domain is not None: - universe_domain = client_universe_domain - elif universe_domain_env is not None: - universe_domain = universe_domain_env - if len(universe_domain.strip()) == 0: - raise ValueError("Universe Domain cannot be an empty string.") - return universe_domain - def _validate_universe_domain(self): """Validates client's and credentials' universe domains are consistent. @@ -697,13 +558,15 @@ def __init__( universe_domain_opt = getattr(self._client_options, "universe_domain", None) self._use_client_cert, self._use_mtls_endpoint, self._universe_domain_env = ( - VizierServiceClient._read_environment_variables() + read_environment_variables() ) self._client_cert_source = VizierServiceClient._get_client_cert_source( self._client_options.client_cert_source, self._use_client_cert ) - self._universe_domain = VizierServiceClient._get_universe_domain( - universe_domain_opt, self._universe_domain_env + self._universe_domain = get_universe_domain( + universe_domain_opt, + self._universe_domain_env, + default_universe=VizierServiceClient._DEFAULT_UNIVERSE, ) self._api_endpoint: str = "" # updated below, depending on `transport` @@ -739,14 +602,14 @@ def __init__( self._transport = cast(VizierServiceTransport, transport) self._api_endpoint = self._transport.host - self._api_endpoint = ( - self._api_endpoint - or VizierServiceClient._get_api_endpoint( - self._client_options.api_endpoint, - self._client_cert_source, - self._universe_domain, - self._use_mtls_endpoint, - ) + self._api_endpoint = self._api_endpoint or get_api_endpoint( + api_override=self._client_options.api_endpoint, + universe_domain=self._universe_domain, + default_universe=VizierServiceClient._DEFAULT_UNIVERSE, + default_mtls_endpoint=VizierServiceClient.DEFAULT_MTLS_ENDPOINT, + default_endpoint_template=VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE, + use_mtls=self._use_mtls_endpoint == "always" + or (self._use_mtls_endpoint == "auto" and self._client_cert_source), ) if not transport_provided: diff --git a/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest.py b/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest.py index 59b423fc26..8c3d6a4140 100644 --- a/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest.py +++ b/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest.py @@ -23,6 +23,7 @@ from google.api_core import rest_helpers from google.api_core import rest_streaming from google.api_core import gapic_v1 +from google.cloud.aiplatform_v1beta1._compat import transcode_request import google.protobuf from google.protobuf import json_format @@ -3463,21 +3464,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_http_options() ) - request, metadata = self._interceptor.pre_add_trial_measurement( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseAddTrialMeasurement, + "_BaseAddTrialMeasurement__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3622,21 +3620,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_http_options() ) - request, metadata = self._interceptor.pre_check_trial_early_stopping_state( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState, + "_BaseCheckTrialEarlyStoppingState__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3783,19 +3778,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCompleteTrial._get_http_options() ) - request, metadata = self._interceptor.pre_complete_trial(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCompleteTrial, + "_BaseCompleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3935,23 +3927,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateStudy._get_http_options() ) - request, metadata = self._interceptor.pre_create_study(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateStudy, + "_BaseCreateStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4096,23 +4081,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateTrial._get_http_options() ) - request, metadata = self._interceptor.pre_create_trial(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateTrial._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateTrial, + "_BaseCreateTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4247,17 +4225,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteStudy._get_http_options() ) - request, metadata = self._interceptor.pre_delete_study(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteStudy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteStudy, + "_BaseDeleteStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4357,17 +4334,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteTrial._get_http_options() ) - request, metadata = self._interceptor.pre_delete_trial(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteTrial._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteTrial, + "_BaseDeleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4471,19 +4447,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetStudy._get_http_options() ) - request, metadata = self._interceptor.pre_get_study(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetStudy, + "_BaseGetStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4626,19 +4599,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetTrial._get_http_options() ) - request, metadata = self._interceptor.pre_get_trial(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetTrial, + "_BaseGetTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4779,21 +4749,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_http_options() ) - request, metadata = self._interceptor.pre_list_optimal_trials( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOptimalTrials, + "_BaseListOptimalTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -4936,17 +4903,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListStudies._get_http_options() ) - request, metadata = self._interceptor.pre_list_studies(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseListStudies._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListStudies._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListStudies, + "_BaseListStudies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5088,19 +5054,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListTrials._get_http_options() ) - request, metadata = self._interceptor.pre_list_trials(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListTrials, + "_BaseListTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5241,23 +5204,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseLookupStudy._get_http_options() ) - request, metadata = self._interceptor.pre_lookup_study(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseLookupStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseLookupStudy, + "_BaseLookupStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5402,25 +5358,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseStopTrial._get_http_options() ) - request, metadata = self._interceptor.pre_stop_trial(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseStopTrial, + "_BaseStopTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5563,19 +5510,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSuggestTrials._get_http_options() ) - request, metadata = self._interceptor.pre_suggest_trials(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSuggestTrials, + "_BaseSuggestTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5837,17 +5781,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = self._interceptor.pre_get_location(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -5980,15 +5923,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = self._interceptor.pre_list_locations(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6122,19 +6066,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_get_iam_policy(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6269,19 +6210,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = self._interceptor.pre_set_iam_policy(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6416,21 +6354,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6561,17 +6496,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6676,17 +6612,18 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6794,15 +6731,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = self._interceptor.pre_get_operation(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6935,15 +6873,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = self._interceptor.pre_list_operations(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7076,15 +7015,16 @@ def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = self._interceptor.pre_wait_operation(request, metadata) - transcoded_request = _BaseVizierServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_asyncio.py b/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_asyncio.py index 923821475b..c1308a57bb 100644 --- a/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_asyncio.py +++ b/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_asyncio.py @@ -37,6 +37,8 @@ from google.api_core import retry_async as retries from google.api_core import rest_helpers from google.api_core import rest_streaming_async # type: ignore +from google.cloud.aiplatform_v1beta1._compat import transcode_request + import google.protobuf from google.protobuf import json_format @@ -1366,21 +1368,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_http_options() ) - request, metadata = await self._interceptor.pre_add_trial_measurement( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseAddTrialMeasurement, + "_BaseAddTrialMeasurement__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1531,23 +1530,20 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_http_options() ) - request, metadata = ( await self._interceptor.pre_check_trial_early_stopping_state( request, metadata ) ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState, + "_BaseCheckTrialEarlyStoppingState__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1700,21 +1696,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCompleteTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_complete_trial( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCompleteTrial._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCompleteTrial, + "_BaseCompleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -1862,25 +1855,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_create_study( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateStudy, + "_BaseCreateStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2031,25 +2017,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCreateTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_create_trial( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCreateTrial._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCreateTrial, + "_BaseCreateTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2190,19 +2169,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_study( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteStudy._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteStudy, + "_BaseDeleteStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2308,19 +2286,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_trial( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteTrial._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseDeleteTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteTrial, + "_BaseDeleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2430,19 +2407,16 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_study(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetStudy, + "_BaseGetStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2591,19 +2565,16 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_get_trial(request, metadata) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetTrial, + "_BaseGetTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2751,21 +2722,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_http_options() ) - request, metadata = await self._interceptor.pre_list_optimal_trials( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOptimalTrials, + "_BaseListOptimalTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -2916,19 +2884,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListStudies._get_http_options() ) - request, metadata = await self._interceptor.pre_list_studies( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListStudies._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListStudies._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListStudies, + "_BaseListStudies__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3076,21 +3043,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListTrials._get_http_options() ) - request, metadata = await self._interceptor.pre_list_trials( request, metadata ) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_transcoded_request( - http_options, request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseListTrials._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListTrials, + "_BaseListTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3237,25 +3201,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseLookupStudy._get_http_options() ) - request, metadata = await self._interceptor.pre_lookup_study( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseLookupStudy._get_transcoded_request( - http_options, request - ) - - body = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseLookupStudy, + "_BaseLookupStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3406,27 +3363,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseStopTrial._get_http_options() ) - request, metadata = await self._interceptor.pre_stop_trial( request, metadata ) - transcoded_request = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_transcoded_request( - http_options, request - ) - ) - - body = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_request_body_json( - transcoded_request - ) - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseStopTrial._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseStopTrial, + "_BaseStopTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -3575,21 +3523,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSuggestTrials._get_http_options() ) - request, metadata = await self._interceptor.pre_suggest_trials( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSuggestTrials._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSuggestTrials, + "_BaseSuggestTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=True, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6056,19 +6001,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetLocation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_location( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetLocation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = ( - _BaseVizierServiceRestTransport._BaseGetLocation._get_query_params_json( - transcoded_request - ) + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetLocation, + "_BaseGetLocation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6207,17 +6151,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListLocations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_locations( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListLocations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListLocations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListLocations, + "_BaseListLocations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6359,21 +6304,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_get_iam_policy( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetIamPolicy, + "_BaseGetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6516,21 +6458,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_http_options() ) - request, metadata = await self._interceptor.pre_set_iam_policy( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseSetIamPolicy._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseSetIamPolicy, + "_BaseSetIamPolicy__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6674,21 +6613,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_http_options() ) - request, metadata = await self._interceptor.pre_test_iam_permissions( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_transcoded_request( - http_options, request - ) - - body = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_request_body_json( - transcoded_request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseTestIamPermissions._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseTestIamPermissions, + "_BaseTestIamPermissions__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6827,17 +6763,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseCancelOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_cancel_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseCancelOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseCancelOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseCancelOperation, + "_BaseCancelOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -6950,17 +6887,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseDeleteOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_delete_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseDeleteOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseDeleteOperation, + "_BaseDeleteOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7076,17 +7014,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseGetOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_get_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseGetOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseGetOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseGetOperation, + "_BaseGetOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7227,17 +7166,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseListOperations._get_http_options() ) - request, metadata = await self._interceptor.pre_list_operations( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseListOperations._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseListOperations._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseListOperations, + "_BaseListOperations__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( @@ -7378,17 +7318,18 @@ async def __call__( http_options = ( _BaseVizierServiceRestTransport._BaseWaitOperation._get_http_options() ) - request, metadata = await self._interceptor.pre_wait_operation( request, metadata ) - transcoded_request = _BaseVizierServiceRestTransport._BaseWaitOperation._get_transcoded_request( - http_options, request - ) - - # Jsonify the query params - query_params = _BaseVizierServiceRestTransport._BaseWaitOperation._get_query_params_json( - transcoded_request + transcoded_request, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=getattr( + _BaseVizierServiceRestTransport._BaseWaitOperation, + "_BaseWaitOperation__REQUIRED_FIELDS_DEFAULT_VALUES", + None, + ), + rest_numeric_enums=False, ) if CLIENT_LOGGING_SUPPORTED and _LOGGER.isEnabledFor( diff --git a/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_base.py b/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_base.py index 577af7e2b1..3ce79c2999 100644 --- a/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_base.py +++ b/google/cloud/aiplatform_v1beta1/services/vizier_service/transports/rest_base.py @@ -102,14 +102,6 @@ def __hash__(self): # pragma: NO COVER __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -121,52 +113,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.AddTrialMeasurementRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseAddTrialMeasurement._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCheckTrialEarlyStoppingState: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -178,52 +130,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CheckTrialEarlyStoppingStateRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCheckTrialEarlyStoppingState._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCompleteTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -235,52 +147,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CompleteTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCompleteTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -292,52 +164,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CreateStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCreateStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseCreateTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -349,52 +181,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.CreateTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseCreateTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -405,43 +197,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.DeleteStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseDeleteStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseDeleteTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -452,43 +213,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.DeleteTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseDeleteTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -499,43 +229,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.GetStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseGetStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -546,43 +245,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.GetTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseGetTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListOptimalTrials: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -594,52 +262,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.ListOptimalTrialsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseListOptimalTrials._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListStudies: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -650,43 +278,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.ListStudiesRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseListStudies._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseListTrials: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -697,43 +294,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.ListTrialsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseListTrials._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseLookupStudy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -745,52 +311,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.LookupStudyRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseLookupStudy._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseStopTrial: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -802,52 +328,12 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.StopTrialRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseStopTrial._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseSuggestTrials: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") __REQUIRED_FIELDS_DEFAULT_VALUES: Dict[str, Any] = {} - @classmethod - def _get_unset_required_fields(cls, message_dict): - return { - k: v - for k, v in cls.__REQUIRED_FIELDS_DEFAULT_VALUES.items() - if k not in message_dict - } - @staticmethod def _get_http_options(): http_options: List[Dict[str, str]] = [ @@ -859,38 +345,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - pb_request = vizier_service.SuggestTrialsRequest.pb(request) - transcoded_request = path_template.transcode(http_options, pb_request) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - # Jsonify the request body - - body = json_format.MessageToJson( - transcoded_request["body"], use_integers_for_enums=True - ) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads( - json_format.MessageToJson( - transcoded_request["query_params"], - use_integers_for_enums=True, - ) - ) - query_params.update( - _BaseVizierServiceRestTransport._BaseSuggestTrials._get_unset_required_fields( - query_params - ) - ) - - query_params["$alt"] = "json;enum-encoding=int" - return query_params - class _BaseGetLocation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -909,17 +363,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListLocations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -938,17 +381,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1032,22 +464,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseSetIamPolicy: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1138,22 +554,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseTestIamPermissions: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1229,22 +629,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_request_body_json(transcoded_request): - body = json.dumps(transcoded_request["body"]) - return body - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseCancelOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -1655,17 +1039,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseDeleteOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2120,17 +1493,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseGetOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -2593,17 +1955,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseListOperations: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3062,17 +2413,6 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - class _BaseWaitOperation: def __hash__(self): # pragma: NO COVER return NotImplementedError("__hash__ must be implemented.") @@ -3527,16 +2867,5 @@ def _get_http_options(): ] return http_options - @staticmethod - def _get_transcoded_request(http_options, request): - request_kwargs = json_format.MessageToDict(request) - transcoded_request = path_template.transcode(http_options, **request_kwargs) - return transcoded_request - - @staticmethod - def _get_query_params_json(transcoded_request): - query_params = json.loads(json.dumps(transcoded_request["query_params"])) - return query_params - __all__ = ("_BaseVizierServiceRestTransport",) diff --git a/google/cloud/aiplatform_v1beta1/types/__init__.py b/google/cloud/aiplatform_v1beta1/types/__init__.py index 95a260292c..d3c34401ad 100644 --- a/google/cloud/aiplatform_v1beta1/types/__init__.py +++ b/google/cloud/aiplatform_v1beta1/types/__init__.py @@ -813,7 +813,18 @@ ReadIndexDatapointsResponse, ) from .memory_bank import ( + IntermediateExtractedMemory, Memory, + MemoryBankCustomizationConfig, + MemoryConjunctionFilter, + MemoryFilter, + MemoryGenerationTriggerConfig, + MemoryMetadataValue, + MemoryProfile, + MemoryRevision, + MemoryTopicId, + StructuredMemoryConfig, + MemoryType, ) from .memory_bank_service import ( CreateMemoryOperationMetadata, @@ -2228,7 +2239,18 @@ "FindNeighborsResponse", "ReadIndexDatapointsRequest", "ReadIndexDatapointsResponse", + "IntermediateExtractedMemory", "Memory", + "MemoryBankCustomizationConfig", + "MemoryConjunctionFilter", + "MemoryFilter", + "MemoryGenerationTriggerConfig", + "MemoryMetadataValue", + "MemoryProfile", + "MemoryRevision", + "MemoryTopicId", + "StructuredMemoryConfig", + "MemoryType", "CreateMemoryOperationMetadata", "CreateMemoryRequest", "DeleteMemoryOperationMetadata", diff --git a/google/cloud/aiplatform_v1beta1/types/memory_bank.py b/google/cloud/aiplatform_v1beta1/types/memory_bank.py index df06cce454..d44195f4dd 100644 --- a/google/cloud/aiplatform_v1beta1/types/memory_bank.py +++ b/google/cloud/aiplatform_v1beta1/types/memory_bank.py @@ -19,18 +19,51 @@ import proto # type: ignore +from google.cloud.aiplatform_v1beta1.types import content as gca_content +from google.cloud.aiplatform_v1beta1.types import openapi import google.protobuf.duration_pb2 as duration_pb2 # type: ignore +import google.protobuf.struct_pb2 as struct_pb2 # type: ignore import google.protobuf.timestamp_pb2 as timestamp_pb2 # type: ignore __protobuf__ = proto.module( package="google.cloud.aiplatform.v1beta1", manifest={ + "MemoryType", "Memory", + "MemoryTopicId", + "MemoryBankCustomizationConfig", + "StructuredMemoryConfig", + "MemoryRevision", + "IntermediateExtractedMemory", + "MemoryMetadataValue", + "MemoryConjunctionFilter", + "MemoryFilter", + "MemoryGenerationTriggerConfig", + "MemoryProfile", }, ) +class MemoryType(proto.Enum): + r"""The type of Memory. + + Values: + MEMORY_TYPE_UNSPECIFIED (0): + Represents an unspecified memory type. This + value should not be used. + NATURAL_LANGUAGE_COLLECTION (1): + Indicates belonging to a collection of + natural language memories. + STRUCTURED_PROFILE (3): + Indicates belonging to a structured profile. + """ + + MEMORY_TYPE_UNSPECIFIED = 0 + NATURAL_LANGUAGE_COLLECTION = 1 + STRUCTURED_PROFILE = 3 + + class Memory(proto.Message): r"""A memory. @@ -43,40 +76,107 @@ class Memory(proto.Message): Attributes: expire_time (google.protobuf.timestamp_pb2.Timestamp): - Optional. Timestamp of when this resource is considered - expired. This is *always* provided on output, regardless of - what ``expiration`` was sent on input. + Optional. Represents the timestamp of when this resource is + considered expired. This is *always* provided on output when + ``expiration`` is set on input, regardless of whether + ``expire_time`` or ``ttl`` was provided. This field is a member of `oneof`_ ``expiration``. ttl (google.protobuf.duration_pb2.Duration): - Optional. Input only. The TTL for this - resource. The expiration time is computed: now + - TTL. + Optional. Input only. Represents the TTL for + this resource. The expiration time is computed: + now + TTL. This field is a member of `oneof`_ ``expiration``. + revision_expire_time (google.protobuf.timestamp_pb2.Timestamp): + Optional. Input only. Represents the + timestamp of when the revision is considered + expired. If not set, the memory revision will be + kept until manually deleted. + + This field is a member of `oneof`_ ``revision_expiration``. + revision_ttl (google.protobuf.duration_pb2.Duration): + Optional. Input only. Represents the TTL for + the revision. The expiration time is computed: + now + TTL. + + This field is a member of `oneof`_ ``revision_expiration``. + disable_memory_revisions (bool): + Optional. Input only. Indicates whether no + revision will be created for this request. + + This field is a member of `oneof`_ ``revision_expiration``. name (str): - Identifier. The resource name of the Memory. Format: + Identifier. Represents the resource name of the Memory. + Format: ``projects/{project}/locations/{location}/reasoningEngines/{reasoning_engine}/memories/{memory}`` display_name (str): - Optional. Display name of the Memory. + Optional. Represents the display name of the + Memory. description (str): - Optional. Description of the Memory. + Optional. Represents the description of the + Memory. create_time (google.protobuf.timestamp_pb2.Timestamp): - Output only. Timestamp when this Memory was - created. + Output only. Represents the timestamp when + this Memory was created. update_time (google.protobuf.timestamp_pb2.Timestamp): - Output only. Timestamp when this Memory was - most recently updated. + Output only. Represents the timestamp when + this Memory was most recently updated. fact (str): - Required. Semantic knowledge extracted from - the source content. + Optional. Represents semantic knowledge + extracted from the source content. scope (MutableMapping[str, str]): - Required. Immutable. The scope of the Memory. Memories are - isolated within their scope. The scope is defined when - creating or generating memories. Scope values cannot contain - the wildcard character '\*'. + Required. Immutable. Represents the scope of the Memory. + Memories are isolated within their scope. The scope is + defined when creating or generating memories. Scope values + cannot contain the wildcard character '\*'. + topics (MutableSequence[google.cloud.aiplatform_v1beta1.types.MemoryTopicId]): + Optional. Represents the Topics of the + Memory. + revision_labels (MutableMapping[str, str]): + Optional. Input only. Represents the labels + to apply to the Memory Revision created as a + result of this request. + metadata (MutableMapping[str, google.cloud.aiplatform_v1beta1.types.MemoryMetadataValue]): + Optional. Represents user-provided metadata + for the Memory. This information was provided + when creating, updating, or generating the + Memory. It was not generated by Memory Bank. + memory_type (google.cloud.aiplatform_v1beta1.types.MemoryType): + Optional. Represents the type of the memory. If not set, the + ``NATURAL_LANGUAGE_COLLECTION`` type is used. If + ``STRUCTURED_COLLECTION`` or ``STRUCTURED_PROFILE`` is used, + then ``structured_data`` must be provided. + structured_content (google.cloud.aiplatform_v1beta1.types.Memory.StructuredContent): + Optional. Represents the structured content + of the memory. + context (str): + Optional. Represents the context of the + memory. """ + class StructuredContent(proto.Message): + r"""Represents the structured value of the memory. + + Attributes: + data (google.protobuf.struct_pb2.Struct): + Required. Represents the structured value of + the memory. + schema_id (str): + Required. Represents the schema ID for which + this structured memory belongs to. + """ + + data: struct_pb2.Struct = proto.Field( + proto.MESSAGE, + number=1, + message=struct_pb2.Struct, + ) + schema_id: str = proto.Field( + proto.STRING, + number=2, + ) + expire_time: timestamp_pb2.Timestamp = proto.Field( proto.MESSAGE, number=13, @@ -89,6 +189,23 @@ class Memory(proto.Message): oneof="expiration", message=duration_pb2.Duration, ) + revision_expire_time: timestamp_pb2.Timestamp = proto.Field( + proto.MESSAGE, + number=16, + oneof="revision_expiration", + message=timestamp_pb2.Timestamp, + ) + revision_ttl: duration_pb2.Duration = proto.Field( + proto.MESSAGE, + number=17, + oneof="revision_expiration", + message=duration_pb2.Duration, + ) + disable_memory_revisions: bool = proto.Field( + proto.BOOL, + number=18, + oneof="revision_expiration", + ) name: str = proto.Field( proto.STRING, number=1, @@ -120,6 +237,768 @@ class Memory(proto.Message): proto.STRING, number=11, ) + topics: MutableSequence["MemoryTopicId"] = proto.RepeatedField( + proto.MESSAGE, + number=15, + message="MemoryTopicId", + ) + revision_labels: MutableMapping[str, str] = proto.MapField( + proto.STRING, + proto.STRING, + number=19, + ) + metadata: MutableMapping[str, "MemoryMetadataValue"] = proto.MapField( + proto.STRING, + proto.MESSAGE, + number=21, + message="MemoryMetadataValue", + ) + memory_type: "MemoryType" = proto.Field( + proto.ENUM, + number=22, + enum="MemoryType", + ) + structured_content: StructuredContent = proto.Field( + proto.MESSAGE, + number=24, + message=StructuredContent, + ) + context: str = proto.Field( + proto.STRING, + number=25, + ) + + +class MemoryTopicId(proto.Message): + r"""A memory topic identifier. + This will be used to label a Memory and to restrict which topics + are eligible for generation or retrieval. + + This message has `oneof`_ fields (mutually exclusive fields). + For each oneof, at most one member field can be set at the same time. + Setting any member of the oneof automatically clears all other + members. + + .. _oneof: https://proto-plus-python.readthedocs.io/en/stable/fields.html#oneofs-mutually-exclusive-fields + + Attributes: + custom_memory_topic_label (str): + Optional. Represents the custom memory topic + label. + + This field is a member of `oneof`_ ``topic_id``. + managed_memory_topic (google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic.ManagedTopicEnum): + Optional. Represents the managed memory + topic. + + This field is a member of `oneof`_ ``topic_id``. + """ + + custom_memory_topic_label: str = proto.Field( + proto.STRING, + number=1, + oneof="topic_id", + ) + managed_memory_topic: ( + "MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic.ManagedTopicEnum" + ) = proto.Field( + proto.ENUM, + number=2, + oneof="topic_id", + enum="MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic.ManagedTopicEnum", + ) + + +class MemoryBankCustomizationConfig(proto.Message): + r"""Represents configuration for organizing natural language + memories for a particular scope. + + Attributes: + scope_keys (MutableSequence[str]): + Optional. Represents the scope keys (i.e. 'user_id') for + which to use this config. A request's scope must include all + of the provided keys for the config to be used (order does + not matter). If empty, then the config will be used for all + requests that do not have a more specific config. Only one + default config is allowed per Memory Bank. + memory_topics (MutableSequence[google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.MemoryTopic]): + Optional. Represents topics of information + that should be extracted from conversations and + stored as memories. If not set, then Memory + Bank's default topics will be used. + generate_memories_examples (MutableSequence[google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.GenerateMemoriesExample]): + Optional. Provides examples of how to + generate memories for a particular scope. + enable_third_person_memories (bool): + Optional. Indicates whether the memories will + be generated in the third person (i.e. "The user + generates memories with Memory Bank."). By + default, the memories will be generated in the + first person (i.e. "I generate memories with + Memory Bank.") + consolidation_config (google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.ConsolidationConfig): + Optional. Represents configuration for + customizing how memories are consolidated + together. + disable_natural_language_memories (bool): + Optional. Indicates whether natural language memory + generation should be disabled for all requests. By default, + natural language memory generation is enabled. Set this to + ``true`` when you only want to generate structured memories. + """ + + class MemoryTopic(proto.Message): + r"""A topic of information that should be extracted from + conversations and stored as memories. + + This message has `oneof`_ fields (mutually exclusive fields). + For each oneof, at most one member field can be set at the same time. + Setting any member of the oneof automatically clears all other + members. + + .. _oneof: https://proto-plus-python.readthedocs.io/en/stable/fields.html#oneofs-mutually-exclusive-fields + + Attributes: + custom_memory_topic (google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.MemoryTopic.CustomMemoryTopic): + A custom memory topic defined by the + developer. + + This field is a member of `oneof`_ ``topic_type``. + managed_memory_topic (google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic): + A managed memory topic defined by Memory + Bank. + + This field is a member of `oneof`_ ``topic_type``. + """ + + class CustomMemoryTopic(proto.Message): + r"""A custom memory topic defined by the developer. + + Attributes: + label (str): + Required. Represents the label of the topic. + description (str): + Required. Represents the description of the + memory topic. This should explain what + information should be extracted for this topic. + """ + + label: str = proto.Field( + proto.STRING, + number=1, + ) + description: str = proto.Field( + proto.STRING, + number=2, + ) + + class ManagedMemoryTopic(proto.Message): + r"""A managed memory topic defined by the system. + + Attributes: + managed_topic_enum (google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic.ManagedTopicEnum): + Required. Represents the managed topic. + """ + + class ManagedTopicEnum(proto.Enum): + r"""Represents managed topics. + + Values: + MANAGED_TOPIC_ENUM_UNSPECIFIED (0): + Represents an unspecified topic. This value + should not be used. + USER_PERSONAL_INFO (1): + Represents significant personal information + about the User like first names, relationships, + hobbies, important dates. + USER_PREFERENCES (2): + Represents stated or implied likes, dislikes, + preferred styles, or patterns. + KEY_CONVERSATION_DETAILS (3): + Represents important milestones or + conclusions within the dialogue. + EXPLICIT_INSTRUCTIONS (4): + Represents information that the user + explicitly requested to remember or forget. + """ + + MANAGED_TOPIC_ENUM_UNSPECIFIED = 0 + USER_PERSONAL_INFO = 1 + USER_PREFERENCES = 2 + KEY_CONVERSATION_DETAILS = 3 + EXPLICIT_INSTRUCTIONS = 4 + + managed_topic_enum: "MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic.ManagedTopicEnum" = proto.Field( + proto.ENUM, + number=1, + enum="MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic.ManagedTopicEnum", + ) + + custom_memory_topic: ( + "MemoryBankCustomizationConfig.MemoryTopic.CustomMemoryTopic" + ) = proto.Field( + proto.MESSAGE, + number=3, + oneof="topic_type", + message="MemoryBankCustomizationConfig.MemoryTopic.CustomMemoryTopic", + ) + managed_memory_topic: ( + "MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic" + ) = proto.Field( + proto.MESSAGE, + number=4, + oneof="topic_type", + message="MemoryBankCustomizationConfig.MemoryTopic.ManagedMemoryTopic", + ) + + class GenerateMemoriesExample(proto.Message): + r"""An example of how to generate memories for a particular + scope. + + + .. _oneof: https://proto-plus-python.readthedocs.io/en/stable/fields.html#oneofs-mutually-exclusive-fields + + Attributes: + conversation_source (google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.GenerateMemoriesExample.ConversationSource): + A conversation source for the example. + + This field is a member of `oneof`_ ``source``. + generated_memories (MutableSequence[google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.GenerateMemoriesExample.GeneratedMemory]): + Optional. Represents the memories that are + expected to be generated from the input + conversation. An empty list indicates that no + memories are expected to be generated for the + input conversation. + """ + + class ConversationSource(proto.Message): + r"""A conversation source for the example. This is similar to + ``DirectContentsSource``. + + Attributes: + events (MutableSequence[google.cloud.aiplatform_v1beta1.types.MemoryBankCustomizationConfig.GenerateMemoriesExample.ConversationSource.Event]): + Optional. Represents the input conversation + events for the example. + """ + + class Event(proto.Message): + r"""A single conversation event. + + Attributes: + content (google.cloud.aiplatform_v1beta1.types.Content): + Required. Represents the content of the + event. + """ + + content: gca_content.Content = proto.Field( + proto.MESSAGE, + number=1, + message=gca_content.Content, + ) + + events: MutableSequence[ + "MemoryBankCustomizationConfig.GenerateMemoriesExample.ConversationSource.Event" + ] = proto.RepeatedField( + proto.MESSAGE, + number=1, + message="MemoryBankCustomizationConfig.GenerateMemoriesExample.ConversationSource.Event", + ) + + class GeneratedMemory(proto.Message): + r"""A memory generated by the operation. + + Attributes: + fact (str): + Required. Represents the fact to generate a + memory from. + topics (MutableSequence[google.cloud.aiplatform_v1beta1.types.MemoryTopicId]): + Optional. Represents the list of topics that the memory + should be associated with. For example, use + ``custom_memory_topic_label = "jargon"`` if the extracted + memory is an example of memory extraction for the custom + topic ``jargon``. + """ + + fact: str = proto.Field( + proto.STRING, + number=1, + ) + topics: MutableSequence["MemoryTopicId"] = proto.RepeatedField( + proto.MESSAGE, + number=2, + message="MemoryTopicId", + ) + + conversation_source: ( + "MemoryBankCustomizationConfig.GenerateMemoriesExample.ConversationSource" + ) = proto.Field( + proto.MESSAGE, + number=3, + oneof="source", + message="MemoryBankCustomizationConfig.GenerateMemoriesExample.ConversationSource", + ) + generated_memories: MutableSequence[ + "MemoryBankCustomizationConfig.GenerateMemoriesExample.GeneratedMemory" + ] = proto.RepeatedField( + proto.MESSAGE, + number=4, + message="MemoryBankCustomizationConfig.GenerateMemoriesExample.GeneratedMemory", + ) + + class ConsolidationConfig(proto.Message): + r"""Represents configuration for customizing how memories are + consolidated. + + Attributes: + revisions_per_candidate_count (int): + Optional. Represents the maximum number of + revisions to consider for each candidate memory. + If not set, then the default value (1) will be + used, which means that only the latest revision + will be considered. + """ + + revisions_per_candidate_count: int = proto.Field( + proto.INT32, + number=1, + ) + + scope_keys: MutableSequence[str] = proto.RepeatedField( + proto.STRING, + number=1, + ) + memory_topics: MutableSequence[MemoryTopic] = proto.RepeatedField( + proto.MESSAGE, + number=2, + message=MemoryTopic, + ) + generate_memories_examples: MutableSequence[GenerateMemoriesExample] = ( + proto.RepeatedField( + proto.MESSAGE, + number=3, + message=GenerateMemoriesExample, + ) + ) + enable_third_person_memories: bool = proto.Field( + proto.BOOL, + number=4, + ) + consolidation_config: ConsolidationConfig = proto.Field( + proto.MESSAGE, + number=5, + message=ConsolidationConfig, + ) + disable_natural_language_memories: bool = proto.Field( + proto.BOOL, + number=6, + ) + + +class StructuredMemoryConfig(proto.Message): + r"""Represents configuration for organizing structured memories + for a particular scope. + + Attributes: + scope_keys (MutableSequence[str]): + Optional. Represents the scope keys (i.e. 'user_id') for + which to use this config. A request's scope must include all + of the provided keys for the config to be used (order does + not matter). If empty, then the config will be used for all + requests that do not have a more specific config. Only one + default config is allowed per Memory Bank. + schema_configs (MutableSequence[google.cloud.aiplatform_v1beta1.types.StructuredMemoryConfig.SchemaConfig]): + Optional. Represents configuration of the + structured memories' schemas. + """ + + class SchemaConfig(proto.Message): + r"""Schema configuration for structured memories. + + Attributes: + id (str): + Required. Represents the ID of the schema. + Must be 1-63 characters, start with a lowercase + letter, and consist of lowercase letters, + numbers, and hyphens. + schema (google.cloud.aiplatform_v1beta1.types.Schema): + Required. Represents the OpenAPI schema of the structured + memories. The schema ``type`` cannot be ``ARRAY`` when + ``memory_type`` is ``STRUCTURED_PROFILE``. + memory_type (google.cloud.aiplatform_v1beta1.types.MemoryType): + Optional. Represents the type of the structured memories + associated with the schema. If not set, then + ``STRUCTURED_PROFILE`` will be used. + json_schema (google.protobuf.struct_pb2.Value): + Optional. Represents the JSON Schema of the + structured memories. + """ + + id: str = proto.Field( + proto.STRING, + number=1, + ) + schema: openapi.Schema = proto.Field( + proto.MESSAGE, + number=2, + message=openapi.Schema, + ) + memory_type: "MemoryType" = proto.Field( + proto.ENUM, + number=3, + enum="MemoryType", + ) + json_schema: struct_pb2.Value = proto.Field( + proto.MESSAGE, + number=5, + message=struct_pb2.Value, + ) + + scope_keys: MutableSequence[str] = proto.RepeatedField( + proto.STRING, + number=1, + ) + schema_configs: MutableSequence[SchemaConfig] = proto.RepeatedField( + proto.MESSAGE, + number=2, + message=SchemaConfig, + ) + + +class MemoryRevision(proto.Message): + r"""A revision of a Memory. + + Attributes: + name (str): + Identifier. Represents the resource name of the Memory + Revision. Format: + ``projects/{project}/locations/{location}/reasoningEngines/{reasoning_engine}/memories/{memory}/revisions/{memory_revision}`` + create_time (google.protobuf.timestamp_pb2.Timestamp): + Output only. Represents the timestamp when + this Memory Revision was created. + expire_time (google.protobuf.timestamp_pb2.Timestamp): + Output only. Represents the timestamp of when + this resource is considered expired. + fact (str): + Output only. Represents the fact of the Memory Revision. + This corresponds to the ``fact`` field of the parent Memory + at the time of revision creation. + labels (MutableMapping[str, str]): + Output only. Represents the labels of the Memory Revision. + These labels are applied to the MemoryRevision when it is + created based on + ``GenerateMemoriesRequest.revision_labels``. + extracted_memories (MutableSequence[google.cloud.aiplatform_v1beta1.types.IntermediateExtractedMemory]): + Output only. Represents the extracted + memories from the source content before + consolidation when the memory was updated via + GenerateMemories. This information was used to + modify an existing Memory via Consolidation. + structured_data (google.protobuf.struct_pb2.Struct): + Output only. Represents the structured value + of the memory at the time of revision creation. + context (str): + Output only. Represents the context of the + Memory Revision. The context may include context + from both the historical revisions and the + extracted content. + """ + + name: str = proto.Field( + proto.STRING, + number=1, + ) + create_time: timestamp_pb2.Timestamp = proto.Field( + proto.MESSAGE, + number=2, + message=timestamp_pb2.Timestamp, + ) + expire_time: timestamp_pb2.Timestamp = proto.Field( + proto.MESSAGE, + number=3, + message=timestamp_pb2.Timestamp, + ) + fact: str = proto.Field( + proto.STRING, + number=4, + ) + labels: MutableMapping[str, str] = proto.MapField( + proto.STRING, + proto.STRING, + number=5, + ) + extracted_memories: MutableSequence["IntermediateExtractedMemory"] = ( + proto.RepeatedField( + proto.MESSAGE, + number=6, + message="IntermediateExtractedMemory", + ) + ) + structured_data: struct_pb2.Struct = proto.Field( + proto.MESSAGE, + number=7, + message=struct_pb2.Struct, + ) + context: str = proto.Field( + proto.STRING, + number=8, + ) + + +class IntermediateExtractedMemory(proto.Message): + r"""An extracted memory that is the intermediate result before + consolidation. + + Attributes: + fact (str): + Output only. Represents the fact of the + extracted memory. + structured_data (google.protobuf.struct_pb2.Struct): + Output only. Represents the structured value + of the extracted memory. + context (str): + Output only. Represents the explanation of + why the information was extracted from the + source content. + """ + + fact: str = proto.Field( + proto.STRING, + number=1, + ) + structured_data: struct_pb2.Struct = proto.Field( + proto.MESSAGE, + number=3, + message=struct_pb2.Struct, + ) + context: str = proto.Field( + proto.STRING, + number=4, + ) + + +class MemoryMetadataValue(proto.Message): + r"""Memory metadata. + + This message has `oneof`_ fields (mutually exclusive fields). + For each oneof, at most one member field can be set at the same time. + Setting any member of the oneof automatically clears all other + members. + + .. _oneof: https://proto-plus-python.readthedocs.io/en/stable/fields.html#oneofs-mutually-exclusive-fields + + Attributes: + string_value (str): + Represents a string value. + + This field is a member of `oneof`_ ``value``. + double_value (float): + Represents a double value. + + This field is a member of `oneof`_ ``value``. + bool_value (bool): + Represents a boolean value. + + This field is a member of `oneof`_ ``value``. + timestamp_value (google.protobuf.timestamp_pb2.Timestamp): + Represents a timestamp value. When filtering + on timestamp values, only the seconds field will + be compared. + + This field is a member of `oneof`_ ``value``. + """ + + string_value: str = proto.Field( + proto.STRING, + number=1, + oneof="value", + ) + double_value: float = proto.Field( + proto.DOUBLE, + number=2, + oneof="value", + ) + bool_value: bool = proto.Field( + proto.BOOL, + number=3, + oneof="value", + ) + timestamp_value: timestamp_pb2.Timestamp = proto.Field( + proto.MESSAGE, + number=4, + oneof="value", + message=timestamp_pb2.Timestamp, + ) + + +class MemoryConjunctionFilter(proto.Message): + r"""A conjunction of filters that will be combined using AND + logic. + + Attributes: + filters (MutableSequence[google.cloud.aiplatform_v1beta1.types.MemoryFilter]): + Represents filters that will be combined + using AND logic. + """ + + filters: MutableSequence["MemoryFilter"] = proto.RepeatedField( + proto.MESSAGE, + number=1, + message="MemoryFilter", + ) + + +class MemoryFilter(proto.Message): + r"""Filter to apply when retrieving memories. + + Attributes: + key (str): + Represents the key of the filter. For example, "author" + would apply to ``metadata`` entries with the key "author". + op (google.cloud.aiplatform_v1beta1.types.MemoryFilter.Operator): + Represents the operator to apply to the + filter. If not set, then EQUAL will be used. + value (google.cloud.aiplatform_v1beta1.types.MemoryMetadataValue): + Represents the value to compare to. + negate (bool): + Indicates whether the filter will be negated. + """ + + class Operator(proto.Enum): + r"""Represents the operator to apply to the filter. + + Values: + OPERATOR_UNSPECIFIED (0): + Represents an unspecified operator. Defaults + to EQUAL. + EQUAL (1): + Equal to. + GREATER_THAN (2): + Greater than. + LESS_THAN (3): + Less than. + """ + + OPERATOR_UNSPECIFIED = 0 + EQUAL = 1 + GREATER_THAN = 2 + LESS_THAN = 3 + + key: str = proto.Field( + proto.STRING, + number=1, + ) + op: Operator = proto.Field( + proto.ENUM, + number=2, + enum=Operator, + ) + value: "MemoryMetadataValue" = proto.Field( + proto.MESSAGE, + number=3, + message="MemoryMetadataValue", + ) + negate: bool = proto.Field( + proto.BOOL, + number=4, + ) + + +class MemoryGenerationTriggerConfig(proto.Message): + r"""Represents configuration for triggering generation. + + Attributes: + generation_rule (google.cloud.aiplatform_v1beta1.types.MemoryGenerationTriggerConfig.GenerationTriggerRule): + Optional. Represents the active rule that + determines when to flush the buffer. If not set, + then the stream will be force flushed + immediately. + """ + + class GenerationTriggerRule(proto.Message): + r"""Represents the active rule that determines when to flush the + buffer. + + This message has `oneof`_ fields (mutually exclusive fields). + For each oneof, at most one member field can be set at the same time. + Setting any member of the oneof automatically clears all other + members. + + .. _oneof: https://proto-plus-python.readthedocs.io/en/stable/fields.html#oneofs-mutually-exclusive-fields + + Attributes: + idle_duration (google.protobuf.duration_pb2.Duration): + Optional. Specifies to trigger generation if + the stream is inactive for the specified + duration after the most recent event. The + duration must have a minute-level granularity. + + This field is a member of `oneof`_ ``time_based_condition``. + fixed_interval (google.protobuf.duration_pb2.Duration): + Optional. Specifies to trigger generation at + a fixed interval. The duration must have a + minute-level granularity. + + This field is a member of `oneof`_ ``time_based_condition``. + overlap_event_count (int): + Optional. Re-include the last N + already-processed events in the next window. + + This field is a member of `oneof`_ ``overlap_window``. + event_count (int): + Optional. Specifies to trigger generation + when the event count reaches this limit. + """ + + idle_duration: duration_pb2.Duration = proto.Field( + proto.MESSAGE, + number=1, + oneof="time_based_condition", + message=duration_pb2.Duration, + ) + fixed_interval: duration_pb2.Duration = proto.Field( + proto.MESSAGE, + number=2, + oneof="time_based_condition", + message=duration_pb2.Duration, + ) + overlap_event_count: int = proto.Field( + proto.INT32, + number=5, + oneof="overlap_window", + ) + event_count: int = proto.Field( + proto.INT32, + number=4, + ) + + generation_rule: GenerationTriggerRule = proto.Field( + proto.MESSAGE, + number=1, + message=GenerationTriggerRule, + ) + + +class MemoryProfile(proto.Message): + r"""A memory profile. + + Attributes: + schema_id (str): + Represents the ID of the schema. This ID corresponds to the + ``schema_id`` defined inside the SchemaConfig, under + StructuredMemoryCustomizationConfig. + profile (google.protobuf.struct_pb2.Struct): + Represents the profile data. + """ + + schema_id: str = proto.Field( + proto.STRING, + number=1, + ) + profile: struct_pb2.Struct = proto.Field( + proto.MESSAGE, + number=2, + message=struct_pb2.Struct, + ) __all__ = tuple(sorted(__protobuf__.manifest)) diff --git a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_async.py b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_async.py index caabdb1a0d..71600c01f3 100644 --- a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_async.py +++ b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_async.py @@ -39,12 +39,8 @@ async def sample_create_memory(): client = aiplatform_v1beta1.MemoryBankServiceAsyncClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.CreateMemoryRequest( parent="parent_value", - memory=memory, ) # Make the request diff --git a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_sync.py b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_sync.py index 115c66a712..b15db6399d 100644 --- a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_sync.py +++ b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_create_memory_sync.py @@ -39,12 +39,8 @@ def sample_create_memory(): client = aiplatform_v1beta1.MemoryBankServiceClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.CreateMemoryRequest( parent="parent_value", - memory=memory, ) # Make the request diff --git a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_async.py b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_async.py index 56396a5a9b..a08ca5a92f 100644 --- a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_async.py +++ b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_async.py @@ -39,11 +39,7 @@ async def sample_update_memory(): client = aiplatform_v1beta1.MemoryBankServiceAsyncClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.UpdateMemoryRequest( - memory=memory, ) # Make the request diff --git a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_sync.py b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_sync.py index 0e24629023..aca7315507 100644 --- a/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_sync.py +++ b/samples/generated_samples/aiplatform_v1beta1_generated_memory_bank_service_update_memory_sync.py @@ -39,11 +39,7 @@ def sample_update_memory(): client = aiplatform_v1beta1.MemoryBankServiceClient() # Initialize request argument(s) - memory = aiplatform_v1beta1.Memory() - memory.fact = "fact_value" - request = aiplatform_v1beta1.UpdateMemoryRequest( - memory=memory, ) # Make the request diff --git a/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1.json b/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1.json index 72bd680368..53753f28da 100644 --- a/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1.json +++ b/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1.json @@ -8,7 +8,7 @@ ], "language": "PYTHON", "name": "google-cloud-aiplatform", - "version": "2.2.0" + "version": "0.0.0" }, "snippets": [ { diff --git a/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1beta1.json b/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1beta1.json index c26e8327c9..1211cc13ad 100644 --- a/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1beta1.json +++ b/samples/generated_samples/snippet_metadata_google.cloud.aiplatform.v1beta1.json @@ -8,7 +8,7 @@ ], "language": "PYTHON", "name": "google-cloud-aiplatform", - "version": "2.2.0" + "version": "0.0.0" }, "snippets": [ { @@ -29804,12 +29804,12 @@ "regionTag": "aiplatform_v1beta1_generated_MemoryBankService_CreateMemory_async", "segments": [ { - "end": 59, + "end": 55, "start": 27, "type": "FULL" }, { - "end": 59, + "end": 55, "start": 27, "type": "SHORT" }, @@ -29819,18 +29819,18 @@ "type": "CLIENT_INITIALIZATION" }, { - "end": 49, + "end": 45, "start": 41, "type": "REQUEST_INITIALIZATION" }, { - "end": 56, - "start": 50, + "end": 52, + "start": 46, "type": "REQUEST_EXECUTION" }, { - "end": 60, - "start": 57, + "end": 56, + "start": 53, "type": "RESPONSE_HANDLING" } ], @@ -29892,12 +29892,12 @@ "regionTag": "aiplatform_v1beta1_generated_MemoryBankService_CreateMemory_sync", "segments": [ { - "end": 59, + "end": 55, "start": 27, "type": "FULL" }, { - "end": 59, + "end": 55, "start": 27, "type": "SHORT" }, @@ -29907,18 +29907,18 @@ "type": "CLIENT_INITIALIZATION" }, { - "end": 49, + "end": 45, "start": 41, "type": "REQUEST_INITIALIZATION" }, { - "end": 56, - "start": 50, + "end": 52, + "start": 46, "type": "REQUEST_EXECUTION" }, { - "end": 60, - "start": 57, + "end": 56, + "start": 53, "type": "RESPONSE_HANDLING" } ], @@ -30782,12 +30782,12 @@ "regionTag": "aiplatform_v1beta1_generated_MemoryBankService_UpdateMemory_async", "segments": [ { - "end": 58, + "end": 54, "start": 27, "type": "FULL" }, { - "end": 58, + "end": 54, "start": 27, "type": "SHORT" }, @@ -30797,18 +30797,18 @@ "type": "CLIENT_INITIALIZATION" }, { - "end": 48, + "end": 44, "start": 41, "type": "REQUEST_INITIALIZATION" }, { - "end": 55, - "start": 49, + "end": 51, + "start": 45, "type": "REQUEST_EXECUTION" }, { - "end": 59, - "start": 56, + "end": 55, + "start": 52, "type": "RESPONSE_HANDLING" } ], @@ -30866,12 +30866,12 @@ "regionTag": "aiplatform_v1beta1_generated_MemoryBankService_UpdateMemory_sync", "segments": [ { - "end": 58, + "end": 54, "start": 27, "type": "FULL" }, { - "end": 58, + "end": 54, "start": 27, "type": "SHORT" }, @@ -30881,18 +30881,18 @@ "type": "CLIENT_INITIALIZATION" }, { - "end": 48, + "end": 44, "start": 41, "type": "REQUEST_INITIALIZATION" }, { - "end": 55, - "start": 49, + "end": 51, + "start": 45, "type": "REQUEST_EXECUTION" }, { - "end": 59, - "start": 56, + "end": 55, + "start": 52, "type": "RESPONSE_HANDLING" } ], diff --git a/tests/unit/gapic/aiplatform_v1/test_compat.py b/tests/unit/gapic/aiplatform_v1/test_compat.py new file mode 100644 index 0000000000..eb136fbd77 --- /dev/null +++ b/tests/unit/gapic/aiplatform_v1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform_v1._compat import transcode_request +from google.cloud.aiplatform_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/aiplatform_v1/test_data_foundry_service.py b/tests/unit/gapic/aiplatform_v1/test_data_foundry_service.py index 2c25dc36bb..86f29ff80d 100644 --- a/tests/unit/gapic/aiplatform_v1/test_data_foundry_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_data_foundry_service.py @@ -82,6 +82,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -125,150 +137,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert DataFoundryServiceClient._get_default_mtls_endpoint(None) is None - assert DataFoundryServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert DataFoundryServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert DataFoundryServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert DataFoundryServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert DataFoundryServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert DataFoundryServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert DataFoundryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DataFoundryServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DataFoundryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - DataFoundryServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert DataFoundryServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert DataFoundryServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert DataFoundryServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert DataFoundryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - DataFoundryServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert DataFoundryServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert DataFoundryServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert DataFoundryServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DataFoundryServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DataFoundryServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert DataFoundryServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert DataFoundryServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert DataFoundryServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert DataFoundryServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert DataFoundryServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - DataFoundryServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert DataFoundryServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert DataFoundryServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -282,40 +150,6 @@ def test__get_client_cert_source(): assert DataFoundryServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert DataFoundryServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(DataFoundryServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DataFoundryServiceClient)) -@mock.patch.object(DataFoundryServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DataFoundryServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = DataFoundryServiceClient._DEFAULT_UNIVERSE - default_endpoint = DataFoundryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = DataFoundryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert DataFoundryServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert DataFoundryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == DataFoundryServiceClient.DEFAULT_MTLS_ENDPOINT - assert DataFoundryServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert DataFoundryServiceClient._get_api_endpoint(None, None, default_universe, "always") == DataFoundryServiceClient.DEFAULT_MTLS_ENDPOINT - assert DataFoundryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == DataFoundryServiceClient.DEFAULT_MTLS_ENDPOINT - assert DataFoundryServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert DataFoundryServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - DataFoundryServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert DataFoundryServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert DataFoundryServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert DataFoundryServiceClient._get_universe_domain(None, None) == DataFoundryServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - DataFoundryServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -700,11 +534,12 @@ def test_data_foundry_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -747,11 +582,12 @@ def test_data_foundry_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1256,7 +1092,12 @@ def test_generate_synthetic_data_rest_required_fields(request_type=data_foundry_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_synthetic_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGenerateSyntheticData, + "_BaseGenerateSyntheticData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -1264,9 +1105,6 @@ def test_generate_synthetic_data_rest_required_fields(request_type=data_foundry_ jsonified_request["location"] = 'location_value' jsonified_request["count"] = 553 - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_synthetic_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "location" in jsonified_request assert jsonified_request["location"] == 'location_value' @@ -1318,13 +1156,6 @@ def test_generate_synthetic_data_rest_required_fields(request_type=data_foundry_ assert sorted(expected_params) == sorted(actual_params) -def test_generate_synthetic_data_rest_unset_required_fields(): - transport = transports.DataFoundryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.generate_synthetic_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("location", "count", "outputFieldSpecs", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.DataFoundryServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1/test_dataset_service.py b/tests/unit/gapic/aiplatform_v1/test_dataset_service.py index 83e329843d..57b5457bab 100644 --- a/tests/unit/gapic/aiplatform_v1/test_dataset_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_dataset_service.py @@ -98,6 +98,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -141,150 +153,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert DatasetServiceClient._get_default_mtls_endpoint(None) is None - assert DatasetServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert DatasetServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DatasetServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - DatasetServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert DatasetServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert DatasetServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert DatasetServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - DatasetServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - DatasetServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert DatasetServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -298,40 +166,6 @@ def test__get_client_cert_source(): assert DatasetServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert DatasetServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(DatasetServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DatasetServiceClient)) -@mock.patch.object(DatasetServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DatasetServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = DatasetServiceClient._DEFAULT_UNIVERSE - default_endpoint = DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert DatasetServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert DatasetServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - assert DatasetServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert DatasetServiceClient._get_api_endpoint(None, None, default_universe, "always") == DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - assert DatasetServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - assert DatasetServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert DatasetServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - DatasetServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert DatasetServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert DatasetServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert DatasetServiceClient._get_universe_domain(None, None) == DatasetServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - DatasetServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -716,11 +550,12 @@ def test_dataset_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -763,11 +598,12 @@ def test_dataset_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -8328,16 +8164,18 @@ def test_create_dataset_rest_required_fields(request_type=dataset_service.Create # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDataset, + "_BaseCreateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -8384,13 +8222,6 @@ def test_create_dataset_rest_required_fields(request_type=dataset_service.Create assert sorted(expected_params) == sorted(actual_params) -def test_create_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "dataset", ))) - - def test_create_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8493,17 +8324,20 @@ def test_get_dataset_rest_required_fields(request_type=dataset_service.GetDatase # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDataset, + "_BaseGetDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -8553,13 +8387,6 @@ def test_get_dataset_rest_required_fields(request_type=dataset_service.GetDatase assert sorted(expected_params) == sorted(actual_params) -def test_get_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(("readMask", )) & set(("name", ))) - - def test_get_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8661,15 +8488,18 @@ def test_update_dataset_rest_required_fields(request_type=dataset_service.Update # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateDataset, + "_BaseUpdateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -8718,13 +8548,6 @@ def test_update_dataset_rest_required_fields(request_type=dataset_service.Update assert sorted(expected_params) == sorted(actual_params) -def test_update_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("dataset", "updateMask", ))) - - def test_update_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8829,17 +8652,20 @@ def test_list_datasets_rest_required_fields(request_type=dataset_service.ListDat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_datasets._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDatasets, + "_BaseListDatasets__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_datasets._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8889,13 +8715,6 @@ def test_list_datasets_rest_required_fields(request_type=dataset_service.ListDat assert sorted(expected_params) == sorted(actual_params) -def test_list_datasets_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_datasets._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_datasets_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9067,16 +8886,18 @@ def test_delete_dataset_rest_required_fields(request_type=dataset_service.Delete # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDataset, + "_BaseDeleteDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9122,13 +8943,6 @@ def test_delete_dataset_rest_required_fields(request_type=dataset_service.Delete assert sorted(expected_params) == sorted(actual_params) -def test_delete_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9233,16 +9047,18 @@ def test_import_data_rest_required_fields(request_type=dataset_service.ImportDat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportData, + "_BaseImportData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9289,13 +9105,6 @@ def test_import_data_rest_required_fields(request_type=dataset_service.ImportDat assert sorted(expected_params) == sorted(actual_params) -def test_import_data_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "importConfigs", ))) - - def test_import_data_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9402,16 +9211,18 @@ def test_export_data_rest_required_fields(request_type=dataset_service.ExportDat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportData, + "_BaseExportData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9458,13 +9269,6 @@ def test_export_data_rest_required_fields(request_type=dataset_service.ExportDat assert sorted(expected_params) == sorted(actual_params) -def test_export_data_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "exportConfig", ))) - - def test_export_data_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9571,16 +9375,18 @@ def test_create_dataset_version_rest_required_fields(request_type=dataset_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDatasetVersion, + "_BaseCreateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -9627,13 +9433,6 @@ def test_create_dataset_version_rest_required_fields(request_type=dataset_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "datasetVersion", ))) - - def test_create_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9735,15 +9534,18 @@ def test_update_dataset_version_rest_required_fields(request_type=dataset_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateDatasetVersion, + "_BaseUpdateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset_version._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -9792,13 +9594,6 @@ def test_update_dataset_version_rest_required_fields(request_type=dataset_servic assert sorted(expected_params) == sorted(actual_params) -def test_update_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("datasetVersion", "updateMask", ))) - - def test_update_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9907,16 +9702,18 @@ def test_delete_dataset_version_rest_required_fields(request_type=dataset_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDatasetVersion, + "_BaseDeleteDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9962,13 +9759,6 @@ def test_delete_dataset_version_rest_required_fields(request_type=dataset_servic assert sorted(expected_params) == sorted(actual_params) -def test_delete_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10069,17 +9859,20 @@ def test_get_dataset_version_rest_required_fields(request_type=dataset_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDatasetVersion, + "_BaseGetDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset_version._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -10129,13 +9922,6 @@ def test_get_dataset_version_rest_required_fields(request_type=dataset_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(("readMask", )) & set(("name", ))) - - def test_get_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10238,17 +10024,20 @@ def test_list_dataset_versions_rest_required_fields(request_type=dataset_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_dataset_versions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDatasetVersions, + "_BaseListDatasetVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_dataset_versions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10298,13 +10087,6 @@ def test_list_dataset_versions_rest_required_fields(request_type=dataset_service assert sorted(expected_params) == sorted(actual_params) -def test_list_dataset_versions_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_dataset_versions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_dataset_versions_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10476,16 +10258,18 @@ def test_restore_dataset_version_rest_required_fields(request_type=dataset_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).restore_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRestoreDatasetVersion, + "_BaseRestoreDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).restore_dataset_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10531,13 +10315,6 @@ def test_restore_dataset_version_rest_required_fields(request_type=dataset_servi assert sorted(expected_params) == sorted(actual_params) -def test_restore_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.restore_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_restore_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10638,17 +10415,20 @@ def test_list_data_items_rest_required_fields(request_type=dataset_service.ListD # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_items._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDataItems, + "_BaseListDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_items._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10698,13 +10478,6 @@ def test_list_data_items_rest_required_fields(request_type=dataset_service.ListD assert sorted(expected_params) == sorted(actual_params) -def test_list_data_items_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_data_items._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_data_items_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10872,17 +10645,20 @@ def test_search_data_items_rest_required_fields(request_type=dataset_service.Sea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_data_items._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchDataItems, + "_BaseSearchDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["dataset"] = 'dataset_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_data_items._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("annotation_filters", "annotations_filter", "annotations_limit", "data_item_filter", "data_labeling_job", "field_mask", "order_by", "order_by_annotation", "order_by_data_item", "page_size", "page_token", "saved_query", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("annotationFilters", "annotationsFilter", "annotationsLimit", "dataItemFilter", "dataLabelingJob", "fieldMask", "orderBy", "orderByAnnotation", "orderByDataItem", "pageSize", "pageToken", "savedQuery", )) # verify required fields with non-default values are left alone assert "dataset" in jsonified_request @@ -10932,13 +10708,6 @@ def test_search_data_items_rest_required_fields(request_type=dataset_service.Sea assert sorted(expected_params) == sorted(actual_params) -def test_search_data_items_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_data_items._get_unset_required_fields({}) - assert set(unset_fields) == (set(("annotationFilters", "annotationsFilter", "annotationsLimit", "dataItemFilter", "dataLabelingJob", "fieldMask", "orderBy", "orderByAnnotation", "orderByDataItem", "pageSize", "pageToken", "savedQuery", )) & set(("dataset", ))) - - def test_search_data_items_rest_pager(transport: str = 'rest'): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11052,17 +10821,20 @@ def test_list_saved_queries_rest_required_fields(request_type=dataset_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_saved_queries._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSavedQueries, + "_BaseListSavedQueries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_saved_queries._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11112,13 +10884,6 @@ def test_list_saved_queries_rest_required_fields(request_type=dataset_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_saved_queries_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_saved_queries._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_saved_queries_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11290,16 +11055,18 @@ def test_delete_saved_query_rest_required_fields(request_type=dataset_service.De # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_saved_query._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSavedQuery, + "_BaseDeleteSavedQuery__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_saved_query._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11345,13 +11112,6 @@ def test_delete_saved_query_rest_required_fields(request_type=dataset_service.De assert sorted(expected_params) == sorted(actual_params) -def test_delete_saved_query_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_saved_query._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_saved_query_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11452,17 +11212,20 @@ def test_get_annotation_spec_rest_required_fields(request_type=dataset_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_annotation_spec._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetAnnotationSpec, + "_BaseGetAnnotationSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_annotation_spec._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -11512,13 +11275,6 @@ def test_get_annotation_spec_rest_required_fields(request_type=dataset_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_annotation_spec_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_annotation_spec._get_unset_required_fields({}) - assert set(unset_fields) == (set(("readMask", )) & set(("name", ))) - - def test_get_annotation_spec_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11621,17 +11377,20 @@ def test_list_annotations_rest_required_fields(request_type=dataset_service.List # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_annotations._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListAnnotations, + "_BaseListAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_annotations._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11681,13 +11440,6 @@ def test_list_annotations_rest_required_fields(request_type=dataset_service.List assert sorted(expected_params) == sorted(actual_params) -def test_list_annotations_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_annotations._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_annotations_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_deployment_resource_pool_service.py b/tests/unit/gapic/aiplatform_v1/test_deployment_resource_pool_service.py index 1dd310c9cb..97e9d2f7d5 100644 --- a/tests/unit/gapic/aiplatform_v1/test_deployment_resource_pool_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_deployment_resource_pool_service.py @@ -94,6 +94,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -137,150 +149,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(None) is None - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - DeploymentResourcePoolServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert DeploymentResourcePoolServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - DeploymentResourcePoolServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - DeploymentResourcePoolServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -294,40 +162,6 @@ def test__get_client_cert_source(): assert DeploymentResourcePoolServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert DeploymentResourcePoolServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(DeploymentResourcePoolServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DeploymentResourcePoolServiceClient)) -@mock.patch.object(DeploymentResourcePoolServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DeploymentResourcePoolServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - default_endpoint = DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert DeploymentResourcePoolServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, default_universe, "always") == DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - DeploymentResourcePoolServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert DeploymentResourcePoolServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert DeploymentResourcePoolServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert DeploymentResourcePoolServiceClient._get_universe_domain(None, None) == DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - DeploymentResourcePoolServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -712,11 +546,12 @@ def test_deployment_resource_pool_service_client_get_mtls_endpoint_and_cert_sour for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -759,11 +594,12 @@ def test_deployment_resource_pool_service_client_get_mtls_endpoint_and_cert_sour for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3384,7 +3220,12 @@ def test_create_deployment_resource_pool_rest_required_fields(request_type=deplo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDeploymentResourcePool, + "_BaseCreateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -3392,9 +3233,6 @@ def test_create_deployment_resource_pool_rest_required_fields(request_type=deplo jsonified_request["parent"] = 'parent_value' jsonified_request["deploymentResourcePoolId"] = 'deployment_resource_pool_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_deployment_resource_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3443,13 +3281,6 @@ def test_create_deployment_resource_pool_rest_required_fields(request_type=deplo assert sorted(expected_params) == sorted(actual_params) -def test_create_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "deploymentResourcePool", "deploymentResourcePoolId", ))) - - def test_create_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3554,16 +3385,18 @@ def test_get_deployment_resource_pool_rest_required_fields(request_type=deployme # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDeploymentResourcePool, + "_BaseGetDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_deployment_resource_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3612,13 +3445,6 @@ def test_get_deployment_resource_pool_rest_required_fields(request_type=deployme assert sorted(expected_params) == sorted(actual_params) -def test_get_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3721,17 +3547,20 @@ def test_list_deployment_resource_pools_rest_required_fields(request_type=deploy # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_deployment_resource_pools._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDeploymentResourcePools, + "_BaseListDeploymentResourcePools__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_deployment_resource_pools._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3781,13 +3610,6 @@ def test_list_deployment_resource_pools_rest_required_fields(request_type=deploy assert sorted(expected_params) == sorted(actual_params) -def test_list_deployment_resource_pools_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_deployment_resource_pools._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_deployment_resource_pools_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3958,15 +3780,18 @@ def test_update_deployment_resource_pool_rest_required_fields(request_type=deplo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateDeploymentResourcePool, + "_BaseUpdateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_deployment_resource_pool._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4012,13 +3837,6 @@ def test_update_deployment_resource_pool_rest_required_fields(request_type=deplo assert sorted(expected_params) == sorted(actual_params) -def test_update_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("deploymentResourcePool", "updateMask", ))) - - def test_update_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4125,16 +3943,18 @@ def test_delete_deployment_resource_pool_rest_required_fields(request_type=deplo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDeploymentResourcePool, + "_BaseDeleteDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_deployment_resource_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4180,13 +4000,6 @@ def test_delete_deployment_resource_pool_rest_required_fields(request_type=deplo assert sorted(expected_params) == sorted(actual_params) -def test_delete_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4287,17 +4100,20 @@ def test_query_deployed_models_rest_required_fields(request_type=deployment_reso # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_deployed_models._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryDeployedModels, + "_BaseQueryDeployedModels__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["deploymentResourcePool"] = 'deployment_resource_pool_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_deployed_models._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "deploymentResourcePool" in jsonified_request @@ -4347,13 +4163,6 @@ def test_query_deployed_models_rest_required_fields(request_type=deployment_reso assert sorted(expected_params) == sorted(actual_params) -def test_query_deployed_models_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_deployed_models._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("deploymentResourcePool", ))) - - def test_query_deployed_models_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_endpoint_service.py b/tests/unit/gapic/aiplatform_v1/test_endpoint_service.py index a8c50c3784..65c25ffe43 100644 --- a/tests/unit/gapic/aiplatform_v1/test_endpoint_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_endpoint_service.py @@ -99,6 +99,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -142,150 +154,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert EndpointServiceClient._get_default_mtls_endpoint(None) is None - assert EndpointServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert EndpointServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EndpointServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - EndpointServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert EndpointServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert EndpointServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert EndpointServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - EndpointServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - EndpointServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert EndpointServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -299,40 +167,6 @@ def test__get_client_cert_source(): assert EndpointServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert EndpointServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(EndpointServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EndpointServiceClient)) -@mock.patch.object(EndpointServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EndpointServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = EndpointServiceClient._DEFAULT_UNIVERSE - default_endpoint = EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert EndpointServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert EndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert EndpointServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert EndpointServiceClient._get_api_endpoint(None, None, default_universe, "always") == EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert EndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert EndpointServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert EndpointServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - EndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert EndpointServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert EndpointServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert EndpointServiceClient._get_universe_domain(None, None) == EndpointServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - EndpointServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -717,11 +551,12 @@ def test_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -764,11 +599,12 @@ def test_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -4265,17 +4101,20 @@ def test_create_endpoint_rest_required_fields(request_type=endpoint_service.Crea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateEndpoint, + "_BaseCreateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_endpoint._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("endpoint_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("endpointId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4323,13 +4162,6 @@ def test_create_endpoint_rest_required_fields(request_type=endpoint_service.Crea assert sorted(expected_params) == sorted(actual_params) -def test_create_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(("endpointId", )) & set(("parent", "endpoint", ))) - - def test_create_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4434,16 +4266,18 @@ def test_get_endpoint_rest_required_fields(request_type=endpoint_service.GetEndp # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetEndpoint, + "_BaseGetEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4492,13 +4326,6 @@ def test_get_endpoint_rest_required_fields(request_type=endpoint_service.GetEndp assert sorted(expected_params) == sorted(actual_params) -def test_get_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4601,17 +4428,20 @@ def test_list_endpoints_rest_required_fields(request_type=endpoint_service.ListE # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_endpoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListEndpoints, + "_BaseListEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_endpoints._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4661,13 +4491,6 @@ def test_list_endpoints_rest_required_fields(request_type=endpoint_service.ListE assert sorted(expected_params) == sorted(actual_params) -def test_list_endpoints_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_endpoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_endpoints_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4834,15 +4657,18 @@ def test_update_endpoint_rest_required_fields(request_type=endpoint_service.Upda # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateEndpoint, + "_BaseUpdateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4891,13 +4717,6 @@ def test_update_endpoint_rest_required_fields(request_type=endpoint_service.Upda assert sorted(expected_params) == sorted(actual_params) -def test_update_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("endpoint", "updateMask", ))) - - def test_update_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5005,14 +4824,16 @@ def test_update_endpoint_long_running_rest_required_fields(request_type=endpoint # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint_long_running._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateEndpointLongRunning, + "_BaseUpdateEndpointLongRunning__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint_long_running._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone client = EndpointServiceClient( @@ -5057,13 +4878,6 @@ def test_update_endpoint_long_running_rest_required_fields(request_type=endpoint assert sorted(expected_params) == sorted(actual_params) -def test_update_endpoint_long_running_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_endpoint_long_running._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_update_endpoint_long_running_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5168,16 +4982,18 @@ def test_delete_endpoint_rest_required_fields(request_type=endpoint_service.Dele # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteEndpoint, + "_BaseDeleteEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5223,13 +5039,6 @@ def test_delete_endpoint_rest_required_fields(request_type=endpoint_service.Dele assert sorted(expected_params) == sorted(actual_params) -def test_delete_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5334,16 +5143,18 @@ def test_deploy_model_rest_required_fields(request_type=endpoint_service.DeployM # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeployModel, + "_BaseDeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5390,13 +5201,6 @@ def test_deploy_model_rest_required_fields(request_type=endpoint_service.DeployM assert sorted(expected_params) == sorted(actual_params) -def test_deploy_model_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.deploy_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "deployedModel", ))) - - def test_deploy_model_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5506,7 +5310,12 @@ def test_undeploy_model_rest_required_fields(request_type=endpoint_service.Undep # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUndeployModel, + "_BaseUndeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -5514,9 +5323,6 @@ def test_undeploy_model_rest_required_fields(request_type=endpoint_service.Undep jsonified_request["endpoint"] = 'endpoint_value' jsonified_request["deployedModelId"] = 'deployed_model_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5565,13 +5371,6 @@ def test_undeploy_model_rest_required_fields(request_type=endpoint_service.Undep assert sorted(expected_params) == sorted(actual_params) -def test_undeploy_model_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.undeploy_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "deployedModelId", ))) - - def test_undeploy_model_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5680,16 +5479,18 @@ def test_mutate_deployed_model_rest_required_fields(request_type=endpoint_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseMutateDeployedModel, + "_BaseMutateDeployedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5736,13 +5537,6 @@ def test_mutate_deployed_model_rest_required_fields(request_type=endpoint_servic assert sorted(expected_params) == sorted(actual_params) -def test_mutate_deployed_model_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.mutate_deployed_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "deployedModel", "updateMask", ))) - - def test_mutate_deployed_model_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_evaluation_service.py b/tests/unit/gapic/aiplatform_v1/test_evaluation_service.py index 3c00841b60..4ddf2dd117 100644 --- a/tests/unit/gapic/aiplatform_v1/test_evaluation_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_evaluation_service.py @@ -82,6 +82,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -125,150 +137,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert EvaluationServiceClient._get_default_mtls_endpoint(None) is None - assert EvaluationServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert EvaluationServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EvaluationServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - EvaluationServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert EvaluationServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - EvaluationServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - EvaluationServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert EvaluationServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -282,40 +150,6 @@ def test__get_client_cert_source(): assert EvaluationServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert EvaluationServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(EvaluationServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EvaluationServiceClient)) -@mock.patch.object(EvaluationServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EvaluationServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = EvaluationServiceClient._DEFAULT_UNIVERSE - default_endpoint = EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert EvaluationServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert EvaluationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - assert EvaluationServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert EvaluationServiceClient._get_api_endpoint(None, None, default_universe, "always") == EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - assert EvaluationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - assert EvaluationServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert EvaluationServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - EvaluationServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert EvaluationServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert EvaluationServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert EvaluationServiceClient._get_universe_domain(None, None) == EvaluationServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - EvaluationServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -700,11 +534,12 @@ def test_evaluation_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -747,11 +582,12 @@ def test_evaluation_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1255,16 +1091,18 @@ def test_evaluate_instances_rest_required_fields(request_type=evaluation_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).evaluate_instances._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseEvaluateInstances, + "_BaseEvaluateInstances__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["location"] = 'location_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).evaluate_instances._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "location" in jsonified_request assert jsonified_request["location"] == 'location_value' @@ -1314,13 +1152,6 @@ def test_evaluate_instances_rest_required_fields(request_type=evaluation_service assert sorted(expected_params) == sorted(actual_params) -def test_evaluate_instances_rest_unset_required_fields(): - transport = transports.EvaluationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.evaluate_instances._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("location", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.EvaluationServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1/test_feature_online_store_admin_service.py b/tests/unit/gapic/aiplatform_v1/test_feature_online_store_admin_service.py index 3f04264c1d..d306cc0369 100644 --- a/tests/unit/gapic/aiplatform_v1/test_feature_online_store_admin_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_feature_online_store_admin_service.py @@ -97,6 +97,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -140,150 +152,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(None) is None - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreAdminServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreAdminServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -297,40 +165,6 @@ def test__get_client_cert_source(): assert FeatureOnlineStoreAdminServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeatureOnlineStoreAdminServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeatureOnlineStoreAdminServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreAdminServiceClient)) -@mock.patch.object(FeatureOnlineStoreAdminServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreAdminServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeatureOnlineStoreAdminServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeatureOnlineStoreAdminServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeatureOnlineStoreAdminServiceClient._get_universe_domain(None, None) == FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreAdminServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -715,11 +549,12 @@ def test_feature_online_store_admin_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -762,11 +597,12 @@ def test_feature_online_store_admin_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5848,7 +5684,12 @@ def test_create_feature_online_store_rest_required_fields(request_type=feature_o # verify fields with default values are dropped assert "featureOnlineStoreId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureOnlineStore, + "_BaseCreateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -5858,10 +5699,8 @@ def test_create_feature_online_store_rest_required_fields(request_type=feature_o jsonified_request["parent"] = 'parent_value' jsonified_request["featureOnlineStoreId"] = 'feature_online_store_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_online_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_online_store_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureOnlineStoreId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5915,13 +5754,6 @@ def test_create_feature_online_store_rest_required_fields(request_type=feature_o assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureOnlineStoreId", )) & set(("parent", "featureOnlineStore", "featureOnlineStoreId", ))) - - def test_create_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6026,16 +5858,18 @@ def test_get_feature_online_store_rest_required_fields(request_type=feature_onli # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureOnlineStore, + "_BaseGetFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_online_store._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6084,13 +5918,6 @@ def test_get_feature_online_store_rest_required_fields(request_type=feature_onli assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6193,17 +6020,20 @@ def test_list_feature_online_stores_rest_required_fields(request_type=feature_on # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_online_stores._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureOnlineStores, + "_BaseListFeatureOnlineStores__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_online_stores._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6253,13 +6083,6 @@ def test_list_feature_online_stores_rest_required_fields(request_type=feature_on assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_online_stores_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_online_stores._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_online_stores_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6430,15 +6253,18 @@ def test_update_feature_online_store_rest_required_fields(request_type=feature_o # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeatureOnlineStore, + "_BaseUpdateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_online_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -6484,13 +6310,6 @@ def test_update_feature_online_store_rest_required_fields(request_type=feature_o assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featureOnlineStore", ))) - - def test_update_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6597,17 +6416,20 @@ def test_delete_feature_online_store_rest_required_fields(request_type=feature_o # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureOnlineStore, + "_BaseDeleteFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_online_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -6654,13 +6476,6 @@ def test_delete_feature_online_store_rest_required_fields(request_type=feature_o assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6769,7 +6584,12 @@ def test_create_feature_view_rest_required_fields(request_type=feature_online_st # verify fields with default values are dropped assert "featureViewId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureView, + "_BaseCreateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -6779,10 +6599,8 @@ def test_create_feature_view_rest_required_fields(request_type=feature_online_st jsonified_request["parent"] = 'parent_value' jsonified_request["featureViewId"] = 'feature_view_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_view._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_view_id", "run_sync_immediately", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureViewId", "runSyncImmediately", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6836,13 +6654,6 @@ def test_create_feature_view_rest_required_fields(request_type=feature_online_st assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureViewId", "runSyncImmediately", )) & set(("parent", "featureView", "featureViewId", ))) - - def test_create_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6947,16 +6758,18 @@ def test_get_feature_view_rest_required_fields(request_type=feature_online_store # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureView, + "_BaseGetFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7005,13 +6818,6 @@ def test_get_feature_view_rest_required_fields(request_type=feature_online_store assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7114,17 +6920,20 @@ def test_list_feature_views_rest_required_fields(request_type=feature_online_sto # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_views._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureViews, + "_BaseListFeatureViews__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_views._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -7174,13 +6983,6 @@ def test_list_feature_views_rest_required_fields(request_type=feature_online_sto assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_views_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_views._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_views_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7351,15 +7153,18 @@ def test_update_feature_view_rest_required_fields(request_type=feature_online_st # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeatureView, + "_BaseUpdateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_view._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -7405,13 +7210,6 @@ def test_update_feature_view_rest_required_fields(request_type=feature_online_st assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featureView", ))) - - def test_update_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7518,16 +7316,18 @@ def test_delete_feature_view_rest_required_fields(request_type=feature_online_st # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureView, + "_BaseDeleteFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_view._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7573,13 +7373,6 @@ def test_delete_feature_view_rest_required_fields(request_type=feature_online_st assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7680,16 +7473,18 @@ def test_sync_feature_view_rest_required_fields(request_type=feature_online_stor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).sync_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSyncFeatureView, + "_BaseSyncFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featureView"] = 'feature_view_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).sync_feature_view._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featureView" in jsonified_request assert jsonified_request["featureView"] == 'feature_view_value' @@ -7739,13 +7534,6 @@ def test_sync_feature_view_rest_required_fields(request_type=feature_online_stor assert sorted(expected_params) == sorted(actual_params) -def test_sync_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.sync_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featureView", ))) - - def test_sync_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7848,16 +7636,18 @@ def test_get_feature_view_sync_rest_required_fields(request_type=feature_online_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view_sync._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureViewSync, + "_BaseGetFeatureViewSync__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view_sync._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7906,13 +7696,6 @@ def test_get_feature_view_sync_rest_required_fields(request_type=feature_online_ assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_view_sync_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_view_sync._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_view_sync_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8015,17 +7798,20 @@ def test_list_feature_view_syncs_rest_required_fields(request_type=feature_onlin # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_view_syncs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureViewSyncs, + "_BaseListFeatureViewSyncs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_view_syncs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8075,13 +7861,6 @@ def test_list_feature_view_syncs_rest_required_fields(request_type=feature_onlin assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_view_syncs_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_view_syncs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_view_syncs_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_feature_online_store_service.py b/tests/unit/gapic/aiplatform_v1/test_feature_online_store_service.py index ab3ffe598d..631fb6e723 100644 --- a/tests/unit/gapic/aiplatform_v1/test_feature_online_store_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_feature_online_store_service.py @@ -83,6 +83,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -126,150 +138,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(None) is None - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeatureOnlineStoreServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeatureOnlineStoreServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -283,40 +151,6 @@ def test__get_client_cert_source(): assert FeatureOnlineStoreServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeatureOnlineStoreServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeatureOnlineStoreServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreServiceClient)) -@mock.patch.object(FeatureOnlineStoreServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeatureOnlineStoreServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeatureOnlineStoreServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeatureOnlineStoreServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeatureOnlineStoreServiceClient._get_universe_domain(None, None) == FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -701,11 +535,12 @@ def test_feature_online_store_service_client_get_mtls_endpoint_and_cert_source(c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -748,11 +583,12 @@ def test_feature_online_store_service_client_get_mtls_endpoint_and_cert_source(c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1931,16 +1767,18 @@ def test_fetch_feature_values_rest_required_fields(request_type=feature_online_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseFetchFeatureValues, + "_BaseFetchFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featureView"] = 'feature_view_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featureView" in jsonified_request assert jsonified_request["featureView"] == 'feature_view_value' @@ -1990,13 +1828,6 @@ def test_fetch_feature_values_rest_required_fields(request_type=feature_online_s assert sorted(expected_params) == sorted(actual_params) -def test_fetch_feature_values_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.fetch_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featureView", ))) - - def test_fetch_feature_values_rest_flattened(): client = FeatureOnlineStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2101,16 +1932,18 @@ def test_search_nearest_entities_rest_required_fields(request_type=feature_onlin # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_nearest_entities._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchNearestEntities, + "_BaseSearchNearestEntities__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featureView"] = 'feature_view_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_nearest_entities._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featureView" in jsonified_request assert jsonified_request["featureView"] == 'feature_view_value' @@ -2160,13 +1993,6 @@ def test_search_nearest_entities_rest_required_fields(request_type=feature_onlin assert sorted(expected_params) == sorted(actual_params) -def test_search_nearest_entities_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_nearest_entities._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featureView", "query", ))) - - def test_feature_view_direct_write_rest_unimplemented(): client = FeatureOnlineStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_feature_registry_service.py b/tests/unit/gapic/aiplatform_v1/test_feature_registry_service.py index 14930e53c5..a087ae69ec 100644 --- a/tests/unit/gapic/aiplatform_v1/test_feature_registry_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_feature_registry_service.py @@ -94,6 +94,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -137,150 +149,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(None) is None - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeatureRegistryServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeatureRegistryServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureRegistryServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeatureRegistryServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -294,40 +162,6 @@ def test__get_client_cert_source(): assert FeatureRegistryServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeatureRegistryServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeatureRegistryServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureRegistryServiceClient)) -@mock.patch.object(FeatureRegistryServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureRegistryServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeatureRegistryServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeatureRegistryServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeatureRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeatureRegistryServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeatureRegistryServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeatureRegistryServiceClient._get_universe_domain(None, None) == FeatureRegistryServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeatureRegistryServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -712,11 +546,12 @@ def test_feature_registry_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -759,11 +594,12 @@ def test_feature_registry_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5027,7 +4863,12 @@ def test_create_feature_group_rest_required_fields(request_type=feature_registry # verify fields with default values are dropped assert "featureGroupId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureGroup, + "_BaseCreateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -5037,10 +4878,8 @@ def test_create_feature_group_rest_required_fields(request_type=feature_registry jsonified_request["parent"] = 'parent_value' jsonified_request["featureGroupId"] = 'feature_group_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_group._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_group_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureGroupId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5094,13 +4933,6 @@ def test_create_feature_group_rest_required_fields(request_type=feature_registry assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureGroupId", )) & set(("parent", "featureGroup", "featureGroupId", ))) - - def test_create_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5205,16 +5037,18 @@ def test_get_feature_group_rest_required_fields(request_type=feature_registry_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureGroup, + "_BaseGetFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_group._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5263,13 +5097,6 @@ def test_get_feature_group_rest_required_fields(request_type=feature_registry_se assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5372,17 +5199,20 @@ def test_list_feature_groups_rest_required_fields(request_type=feature_registry_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_groups._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureGroups, + "_BaseListFeatureGroups__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_groups._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5432,13 +5262,6 @@ def test_list_feature_groups_rest_required_fields(request_type=feature_registry_ assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_groups_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_groups._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_groups_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5609,15 +5432,18 @@ def test_update_feature_group_rest_required_fields(request_type=feature_registry # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeatureGroup, + "_BaseUpdateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_group._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -5663,13 +5489,6 @@ def test_update_feature_group_rest_required_fields(request_type=feature_registry assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featureGroup", ))) - - def test_update_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5776,17 +5595,20 @@ def test_delete_feature_group_rest_required_fields(request_type=feature_registry # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureGroup, + "_BaseDeleteFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_group._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -5833,13 +5655,6 @@ def test_delete_feature_group_rest_required_fields(request_type=feature_registry assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5948,7 +5763,12 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C # verify fields with default values are dropped assert "featureId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -5958,10 +5778,8 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C jsonified_request["parent"] = 'parent_value' jsonified_request["featureId"] = 'feature_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6015,13 +5833,6 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureId", )) & set(("parent", "feature", "featureId", ))) - - def test_create_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6130,16 +5941,18 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6186,13 +5999,6 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_features_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_features_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6295,16 +6101,18 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6353,13 +6161,6 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6462,17 +6263,20 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "latest_stats_count", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6522,13 +6326,6 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_features_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_features_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6699,15 +6496,18 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -6753,13 +6553,6 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("feature", ))) - - def test_update_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6866,16 +6659,18 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6921,13 +6716,6 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_featurestore_online_serving_service.py b/tests/unit/gapic/aiplatform_v1/test_featurestore_online_serving_service.py index 905eee02b4..5e5d55faf1 100644 --- a/tests/unit/gapic/aiplatform_v1/test_featurestore_online_serving_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_featurestore_online_serving_service.py @@ -81,6 +81,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -124,150 +136,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(None) is None - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeaturestoreOnlineServingServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreOnlineServingServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeaturestoreOnlineServingServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -281,40 +149,6 @@ def test__get_client_cert_source(): assert FeaturestoreOnlineServingServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeaturestoreOnlineServingServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeaturestoreOnlineServingServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreOnlineServingServiceClient)) -@mock.patch.object(FeaturestoreOnlineServingServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreOnlineServingServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeaturestoreOnlineServingServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeaturestoreOnlineServingServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeaturestoreOnlineServingServiceClient._get_universe_domain(None, None) == FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeaturestoreOnlineServingServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -699,11 +533,12 @@ def test_featurestore_online_serving_service_client_get_mtls_endpoint_and_cert_s for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -746,11 +581,12 @@ def test_featurestore_online_serving_service_client_get_mtls_endpoint_and_cert_s for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1961,7 +1797,12 @@ def test_read_feature_values_rest_required_fields(request_type=featurestore_onli # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadFeatureValues, + "_BaseReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -1969,9 +1810,6 @@ def test_read_feature_values_rest_required_fields(request_type=featurestore_onli jsonified_request["entityType"] = 'entity_type_value' jsonified_request["entityId"] = 'entity_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -2023,13 +1861,6 @@ def test_read_feature_values_rest_required_fields(request_type=featurestore_onli assert sorted(expected_params) == sorted(actual_params) -def test_read_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreOnlineServingServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "entityId", "featureSelector", ))) - - def test_read_feature_values_rest_flattened(): client = FeaturestoreOnlineServingServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2133,7 +1964,12 @@ def test_streaming_read_feature_values_rest_required_fields(request_type=feature # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).streaming_read_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamingReadFeatureValues, + "_BaseStreamingReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -2141,9 +1977,6 @@ def test_streaming_read_feature_values_rest_required_fields(request_type=feature jsonified_request["entityType"] = 'entity_type_value' jsonified_request["entityIds"] = 'entity_ids_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).streaming_read_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -2198,13 +2031,6 @@ def test_streaming_read_feature_values_rest_required_fields(request_type=feature assert sorted(expected_params) == sorted(actual_params) -def test_streaming_read_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreOnlineServingServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.streaming_read_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "entityIds", "featureSelector", ))) - - def test_streaming_read_feature_values_rest_flattened(): client = FeaturestoreOnlineServingServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2310,16 +2136,18 @@ def test_write_feature_values_rest_required_fields(request_type=featurestore_onl # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseWriteFeatureValues, + "_BaseWriteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -2369,13 +2197,6 @@ def test_write_feature_values_rest_required_fields(request_type=featurestore_onl assert sorted(expected_params) == sorted(actual_params) -def test_write_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreOnlineServingServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.write_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "payloads", ))) - - def test_write_feature_values_rest_flattened(): client = FeaturestoreOnlineServingServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_featurestore_service.py b/tests/unit/gapic/aiplatform_v1/test_featurestore_service.py index f687387099..b5ead01fa2 100644 --- a/tests/unit/gapic/aiplatform_v1/test_featurestore_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_featurestore_service.py @@ -99,6 +99,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -142,150 +154,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeaturestoreServiceClient._get_default_mtls_endpoint(None) is None - assert FeaturestoreServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeaturestoreServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeaturestoreServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeaturestoreServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeaturestoreServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -299,40 +167,6 @@ def test__get_client_cert_source(): assert FeaturestoreServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeaturestoreServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeaturestoreServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreServiceClient)) -@mock.patch.object(FeaturestoreServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeaturestoreServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeaturestoreServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeaturestoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeaturestoreServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeaturestoreServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeaturestoreServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeaturestoreServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeaturestoreServiceClient._get_universe_domain(None, None) == FeaturestoreServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeaturestoreServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -717,11 +551,12 @@ def test_featurestore_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -764,11 +599,12 @@ def test_featurestore_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -8701,7 +8537,12 @@ def test_create_featurestore_rest_required_fields(request_type=featurestore_serv # verify fields with default values are dropped assert "featurestoreId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeaturestore, + "_BaseCreateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -8711,10 +8552,8 @@ def test_create_featurestore_rest_required_fields(request_type=featurestore_serv jsonified_request["parent"] = 'parent_value' jsonified_request["featurestoreId"] = 'featurestore_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_featurestore._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("featurestore_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featurestoreId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8768,13 +8607,6 @@ def test_create_featurestore_rest_required_fields(request_type=featurestore_serv assert sorted(expected_params) == sorted(actual_params) -def test_create_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featurestoreId", )) & set(("parent", "featurestore", "featurestoreId", ))) - - def test_create_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8879,16 +8711,18 @@ def test_get_featurestore_rest_required_fields(request_type=featurestore_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeaturestore, + "_BaseGetFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_featurestore._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8937,13 +8771,6 @@ def test_get_featurestore_rest_required_fields(request_type=featurestore_service assert sorted(expected_params) == sorted(actual_params) -def test_get_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9046,17 +8873,20 @@ def test_list_featurestores_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_featurestores._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeaturestores, + "_BaseListFeaturestores__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_featurestores._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9106,13 +8936,6 @@ def test_list_featurestores_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_list_featurestores_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_featurestores._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_featurestores_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9283,15 +9106,18 @@ def test_update_featurestore_rest_required_fields(request_type=featurestore_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeaturestore, + "_BaseUpdateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_featurestore._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -9337,13 +9163,6 @@ def test_update_featurestore_rest_required_fields(request_type=featurestore_serv assert sorted(expected_params) == sorted(actual_params) -def test_update_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featurestore", ))) - - def test_update_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9450,17 +9269,20 @@ def test_delete_featurestore_rest_required_fields(request_type=featurestore_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeaturestore, + "_BaseDeleteFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_featurestore._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9507,13 +9329,6 @@ def test_delete_featurestore_rest_required_fields(request_type=featurestore_serv assert sorted(expected_params) == sorted(actual_params) -def test_delete_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9622,7 +9437,12 @@ def test_create_entity_type_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped assert "entityTypeId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateEntityType, + "_BaseCreateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -9632,10 +9452,8 @@ def test_create_entity_type_rest_required_fields(request_type=featurestore_servi jsonified_request["parent"] = 'parent_value' jsonified_request["entityTypeId"] = 'entity_type_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_entity_type._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("entity_type_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("entityTypeId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9689,13 +9507,6 @@ def test_create_entity_type_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_create_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(("entityTypeId", )) & set(("parent", "entityTypeId", ))) - - def test_create_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9800,16 +9611,18 @@ def test_get_entity_type_rest_required_fields(request_type=featurestore_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetEntityType, + "_BaseGetEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_entity_type._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9858,13 +9671,6 @@ def test_get_entity_type_rest_required_fields(request_type=featurestore_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9967,17 +9773,20 @@ def test_list_entity_types_rest_required_fields(request_type=featurestore_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_entity_types._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListEntityTypes, + "_BaseListEntityTypes__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_entity_types._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10027,13 +9836,6 @@ def test_list_entity_types_rest_required_fields(request_type=featurestore_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_entity_types_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_entity_types._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_entity_types_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10200,15 +10002,18 @@ def test_update_entity_type_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateEntityType, + "_BaseUpdateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_entity_type._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -10257,13 +10062,6 @@ def test_update_entity_type_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_update_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("entityType", ))) - - def test_update_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10372,17 +10170,20 @@ def test_delete_entity_type_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteEntityType, + "_BaseDeleteEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_entity_type._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -10429,13 +10230,6 @@ def test_delete_entity_type_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_delete_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10544,7 +10338,12 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C # verify fields with default values are dropped assert "featureId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -10554,10 +10353,8 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C jsonified_request["parent"] = 'parent_value' jsonified_request["featureId"] = 'feature_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10611,13 +10408,6 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureId", )) & set(("parent", "feature", "featureId", ))) - - def test_create_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10726,16 +10516,18 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10782,13 +10574,6 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_features_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_features_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10891,16 +10676,18 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10949,13 +10736,6 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11058,17 +10838,20 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "latest_stats_count", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11118,13 +10901,6 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_features_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_features_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11291,15 +11067,18 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -11348,13 +11127,6 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("feature", ))) - - def test_update_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11463,16 +11235,18 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11518,13 +11292,6 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11629,16 +11396,18 @@ def test_import_feature_values_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportFeatureValues, + "_BaseImportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -11685,13 +11454,6 @@ def test_import_feature_values_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_import_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "featureSpecs", ))) - - def test_import_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11796,16 +11558,18 @@ def test_batch_read_feature_values_rest_required_fields(request_type=featurestor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchReadFeatureValues, + "_BaseBatchReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featurestore"] = 'featurestore_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featurestore" in jsonified_request assert jsonified_request["featurestore"] == 'featurestore_value' @@ -11852,13 +11616,6 @@ def test_batch_read_feature_values_rest_required_fields(request_type=featurestor assert sorted(expected_params) == sorted(actual_params) -def test_batch_read_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_read_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featurestore", "destination", "entityTypeSpecs", ))) - - def test_batch_read_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11963,16 +11720,18 @@ def test_export_feature_values_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportFeatureValues, + "_BaseExportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -12019,13 +11778,6 @@ def test_export_feature_values_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_export_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "destination", "featureSelector", ))) - - def test_export_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12130,16 +11882,18 @@ def test_delete_feature_values_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureValues, + "_BaseDeleteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -12186,13 +11940,6 @@ def test_delete_feature_values_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", ))) - - def test_delete_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12293,17 +12040,20 @@ def test_search_features_rest_required_fields(request_type=featurestore_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchFeatures, + "_BaseSearchFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["location"] = 'location_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_features._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", "query", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", "query", )) # verify required fields with non-default values are left alone assert "location" in jsonified_request @@ -12353,13 +12103,6 @@ def test_search_features_rest_required_fields(request_type=featurestore_service. assert sorted(expected_params) == sorted(actual_params) -def test_search_features_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", "query", )) & set(("location", ))) - - def test_search_features_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_gen_ai_cache_service.py b/tests/unit/gapic/aiplatform_v1/test_gen_ai_cache_service.py index fb332325c8..b49eb868fc 100644 --- a/tests/unit/gapic/aiplatform_v1/test_gen_ai_cache_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_gen_ai_cache_service.py @@ -90,6 +90,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -133,150 +145,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert GenAiCacheServiceClient._get_default_mtls_endpoint(None) is None - assert GenAiCacheServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert GenAiCacheServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiCacheServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - GenAiCacheServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert GenAiCacheServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiCacheServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - GenAiCacheServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -290,40 +158,6 @@ def test__get_client_cert_source(): assert GenAiCacheServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert GenAiCacheServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(GenAiCacheServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiCacheServiceClient)) -@mock.patch.object(GenAiCacheServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiCacheServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = GenAiCacheServiceClient._DEFAULT_UNIVERSE - default_endpoint = GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert GenAiCacheServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert GenAiCacheServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiCacheServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert GenAiCacheServiceClient._get_api_endpoint(None, None, default_universe, "always") == GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiCacheServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiCacheServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert GenAiCacheServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiCacheServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert GenAiCacheServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert GenAiCacheServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert GenAiCacheServiceClient._get_universe_domain(None, None) == GenAiCacheServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - GenAiCacheServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -708,11 +542,12 @@ def test_gen_ai_cache_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -755,11 +590,12 @@ def test_gen_ai_cache_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2823,16 +2659,18 @@ def test_create_cached_content_rest_required_fields(request_type=gen_ai_cache_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateCachedContent, + "_BaseCreateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_cached_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2882,13 +2720,6 @@ def test_create_cached_content_rest_required_fields(request_type=gen_ai_cache_se assert sorted(expected_params) == sorted(actual_params) -def test_create_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "cachedContent", ))) - - def test_create_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2993,16 +2824,18 @@ def test_get_cached_content_rest_required_fields(request_type=gen_ai_cache_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetCachedContent, + "_BaseGetCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_cached_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3051,13 +2884,6 @@ def test_get_cached_content_rest_required_fields(request_type=gen_ai_cache_servi assert sorted(expected_params) == sorted(actual_params) -def test_get_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3159,15 +2985,18 @@ def test_update_cached_content_rest_required_fields(request_type=gen_ai_cache_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateCachedContent, + "_BaseUpdateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_cached_content._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3216,13 +3045,6 @@ def test_update_cached_content_rest_required_fields(request_type=gen_ai_cache_se assert sorted(expected_params) == sorted(actual_params) -def test_update_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("cachedContent", "updateMask", ))) - - def test_update_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3327,16 +3149,18 @@ def test_delete_cached_content_rest_required_fields(request_type=gen_ai_cache_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteCachedContent, + "_BaseDeleteCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_cached_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3382,13 +3206,6 @@ def test_delete_cached_content_rest_required_fields(request_type=gen_ai_cache_se assert sorted(expected_params) == sorted(actual_params) -def test_delete_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3489,17 +3306,20 @@ def test_list_cached_contents_rest_required_fields(request_type=gen_ai_cache_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_cached_contents._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListCachedContents, + "_BaseListCachedContents__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_cached_contents._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3549,13 +3369,6 @@ def test_list_cached_contents_rest_required_fields(request_type=gen_ai_cache_ser assert sorted(expected_params) == sorted(actual_params) -def test_list_cached_contents_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_cached_contents._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_cached_contents_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_gen_ai_tuning_service.py b/tests/unit/gapic/aiplatform_v1/test_gen_ai_tuning_service.py index 07d33f0aca..1d208f0caf 100644 --- a/tests/unit/gapic/aiplatform_v1/test_gen_ai_tuning_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_gen_ai_tuning_service.py @@ -97,6 +97,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -140,150 +152,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert GenAiTuningServiceClient._get_default_mtls_endpoint(None) is None - assert GenAiTuningServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert GenAiTuningServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiTuningServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - GenAiTuningServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert GenAiTuningServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiTuningServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - GenAiTuningServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -297,40 +165,6 @@ def test__get_client_cert_source(): assert GenAiTuningServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert GenAiTuningServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(GenAiTuningServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiTuningServiceClient)) -@mock.patch.object(GenAiTuningServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiTuningServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = GenAiTuningServiceClient._DEFAULT_UNIVERSE - default_endpoint = GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert GenAiTuningServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert GenAiTuningServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiTuningServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert GenAiTuningServiceClient._get_api_endpoint(None, None, default_universe, "always") == GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiTuningServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiTuningServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert GenAiTuningServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiTuningServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert GenAiTuningServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert GenAiTuningServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert GenAiTuningServiceClient._get_universe_domain(None, None) == GenAiTuningServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - GenAiTuningServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -715,11 +549,12 @@ def test_gen_ai_tuning_service_client_get_mtls_endpoint_and_cert_source(client_c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -762,11 +597,12 @@ def test_gen_ai_tuning_service_client_get_mtls_endpoint_and_cert_source(client_c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2859,16 +2695,18 @@ def test_create_tuning_job_rest_required_fields(request_type=genai_tuning_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTuningJob, + "_BaseCreateTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2918,13 +2756,6 @@ def test_create_tuning_job_rest_required_fields(request_type=genai_tuning_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_tuning_job_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "tuningJob", ))) - - def test_create_tuning_job_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3029,16 +2860,18 @@ def test_get_tuning_job_rest_required_fields(request_type=genai_tuning_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTuningJob, + "_BaseGetTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3087,13 +2920,6 @@ def test_get_tuning_job_rest_required_fields(request_type=genai_tuning_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_tuning_job_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tuning_job_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3196,17 +3022,20 @@ def test_list_tuning_jobs_rest_required_fields(request_type=genai_tuning_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tuning_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTuningJobs, + "_BaseListTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tuning_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3256,13 +3085,6 @@ def test_list_tuning_jobs_rest_required_fields(request_type=genai_tuning_service assert sorted(expected_params) == sorted(actual_params) -def test_list_tuning_jobs_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tuning_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_tuning_jobs_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3430,16 +3252,18 @@ def test_cancel_tuning_job_rest_required_fields(request_type=genai_tuning_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelTuningJob, + "_BaseCancelTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3486,13 +3310,6 @@ def test_cancel_tuning_job_rest_required_fields(request_type=genai_tuning_servic assert sorted(expected_params) == sorted(actual_params) -def test_cancel_tuning_job_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_tuning_job_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3597,16 +3414,18 @@ def test_rebase_tuned_model_rest_required_fields(request_type=genai_tuning_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).rebase_tuned_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRebaseTunedModel, + "_BaseRebaseTunedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).rebase_tuned_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3653,13 +3472,6 @@ def test_rebase_tuned_model_rest_required_fields(request_type=genai_tuning_servi assert sorted(expected_params) == sorted(actual_params) -def test_rebase_tuned_model_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.rebase_tuned_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "tunedModelRef", ))) - - def test_rebase_tuned_model_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_index_endpoint_service.py b/tests/unit/gapic/aiplatform_v1/test_index_endpoint_service.py index 2c434ed743..3652336fb4 100644 --- a/tests/unit/gapic/aiplatform_v1/test_index_endpoint_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_index_endpoint_service.py @@ -94,6 +94,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -137,150 +149,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert IndexEndpointServiceClient._get_default_mtls_endpoint(None) is None - assert IndexEndpointServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert IndexEndpointServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexEndpointServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - IndexEndpointServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert IndexEndpointServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexEndpointServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - IndexEndpointServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -294,40 +162,6 @@ def test__get_client_cert_source(): assert IndexEndpointServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert IndexEndpointServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(IndexEndpointServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexEndpointServiceClient)) -@mock.patch.object(IndexEndpointServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexEndpointServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = IndexEndpointServiceClient._DEFAULT_UNIVERSE - default_endpoint = IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert IndexEndpointServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert IndexEndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexEndpointServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert IndexEndpointServiceClient._get_api_endpoint(None, None, default_universe, "always") == IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexEndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexEndpointServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert IndexEndpointServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexEndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert IndexEndpointServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert IndexEndpointServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert IndexEndpointServiceClient._get_universe_domain(None, None) == IndexEndpointServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - IndexEndpointServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -712,11 +546,12 @@ def test_index_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_ for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -759,11 +594,12 @@ def test_index_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_ for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3885,16 +3721,18 @@ def test_create_index_endpoint_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateIndexEndpoint, + "_BaseCreateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3941,13 +3779,6 @@ def test_create_index_endpoint_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_create_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "indexEndpoint", ))) - - def test_create_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4050,16 +3881,18 @@ def test_get_index_endpoint_rest_required_fields(request_type=index_endpoint_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetIndexEndpoint, + "_BaseGetIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4108,13 +3941,6 @@ def test_get_index_endpoint_rest_required_fields(request_type=index_endpoint_ser assert sorted(expected_params) == sorted(actual_params) -def test_get_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4217,17 +4043,20 @@ def test_list_index_endpoints_rest_required_fields(request_type=index_endpoint_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_index_endpoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListIndexEndpoints, + "_BaseListIndexEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_index_endpoints._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4277,13 +4106,6 @@ def test_list_index_endpoints_rest_required_fields(request_type=index_endpoint_s assert sorted(expected_params) == sorted(actual_params) -def test_list_index_endpoints_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_index_endpoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_index_endpoints_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4450,15 +4272,18 @@ def test_update_index_endpoint_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateIndexEndpoint, + "_BaseUpdateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index_endpoint._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4507,13 +4332,6 @@ def test_update_index_endpoint_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_update_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("indexEndpoint", "updateMask", ))) - - def test_update_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4622,16 +4440,18 @@ def test_delete_index_endpoint_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteIndexEndpoint, + "_BaseDeleteIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4677,13 +4497,6 @@ def test_delete_index_endpoint_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_delete_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4788,16 +4601,18 @@ def test_deploy_index_rest_required_fields(request_type=index_endpoint_service.D # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeployIndex, + "_BaseDeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -4844,13 +4659,6 @@ def test_deploy_index_rest_required_fields(request_type=index_endpoint_service.D assert sorted(expected_params) == sorted(actual_params) -def test_deploy_index_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.deploy_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", "deployedIndex", ))) - - def test_deploy_index_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4958,7 +4766,12 @@ def test_undeploy_index_rest_required_fields(request_type=index_endpoint_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUndeployIndex, + "_BaseUndeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -4966,9 +4779,6 @@ def test_undeploy_index_rest_required_fields(request_type=index_endpoint_service jsonified_request["indexEndpoint"] = 'index_endpoint_value' jsonified_request["deployedIndexId"] = 'deployed_index_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -5017,13 +4827,6 @@ def test_undeploy_index_rest_required_fields(request_type=index_endpoint_service assert sorted(expected_params) == sorted(actual_params) -def test_undeploy_index_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.undeploy_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", "deployedIndexId", ))) - - def test_undeploy_index_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5130,16 +4933,18 @@ def test_mutate_deployed_index_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseMutateDeployedIndex, + "_BaseMutateDeployedIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -5186,13 +4991,6 @@ def test_mutate_deployed_index_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_mutate_deployed_index_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.mutate_deployed_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", "deployedIndex", ))) - - def test_mutate_deployed_index_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_index_service.py b/tests/unit/gapic/aiplatform_v1/test_index_service.py index 813f243c33..5453f438ba 100644 --- a/tests/unit/gapic/aiplatform_v1/test_index_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_index_service.py @@ -92,6 +92,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -135,150 +147,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert IndexServiceClient._get_default_mtls_endpoint(None) is None - assert IndexServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert IndexServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert IndexServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - IndexServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert IndexServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert IndexServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert IndexServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert IndexServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert IndexServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - IndexServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert IndexServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -292,40 +160,6 @@ def test__get_client_cert_source(): assert IndexServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert IndexServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(IndexServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexServiceClient)) -@mock.patch.object(IndexServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = IndexServiceClient._DEFAULT_UNIVERSE - default_endpoint = IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert IndexServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert IndexServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == IndexServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert IndexServiceClient._get_api_endpoint(None, None, default_universe, "always") == IndexServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == IndexServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert IndexServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert IndexServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert IndexServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert IndexServiceClient._get_universe_domain(None, None) == IndexServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - IndexServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -710,11 +544,12 @@ def test_index_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -757,11 +592,12 @@ def test_index_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3310,16 +3146,18 @@ def test_create_index_rest_required_fields(request_type=index_service.CreateInde # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateIndex, + "_BaseCreateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3366,13 +3204,6 @@ def test_create_index_rest_required_fields(request_type=index_service.CreateInde assert sorted(expected_params) == sorted(actual_params) -def test_create_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "index", ))) - - def test_create_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3475,16 +3306,18 @@ def test_get_index_rest_required_fields(request_type=index_service.GetIndexReque # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetIndex, + "_BaseGetIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3533,13 +3366,6 @@ def test_get_index_rest_required_fields(request_type=index_service.GetIndexReque assert sorted(expected_params) == sorted(actual_params) -def test_get_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3642,17 +3468,20 @@ def test_list_indexes_rest_required_fields(request_type=index_service.ListIndexe # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_indexes._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListIndexes, + "_BaseListIndexes__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_indexes._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3702,13 +3531,6 @@ def test_list_indexes_rest_required_fields(request_type=index_service.ListIndexe assert sorted(expected_params) == sorted(actual_params) -def test_list_indexes_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_indexes._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_indexes_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3879,15 +3701,18 @@ def test_update_index_rest_required_fields(request_type=index_service.UpdateInde # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateIndex, + "_BaseUpdateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3933,13 +3758,6 @@ def test_update_index_rest_required_fields(request_type=index_service.UpdateInde assert sorted(expected_params) == sorted(actual_params) -def test_update_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("index", ))) - - def test_update_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4046,16 +3864,18 @@ def test_delete_index_rest_required_fields(request_type=index_service.DeleteInde # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteIndex, + "_BaseDeleteIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4101,13 +3921,6 @@ def test_delete_index_rest_required_fields(request_type=index_service.DeleteInde assert sorted(expected_params) == sorted(actual_params) -def test_delete_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4208,16 +4021,18 @@ def test_upsert_datapoints_rest_required_fields(request_type=index_service.Upser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upsert_datapoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpsertDatapoints, + "_BaseUpsertDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["index"] = 'index_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upsert_datapoints._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "index" in jsonified_request assert jsonified_request["index"] == 'index_value' @@ -4267,13 +4082,6 @@ def test_upsert_datapoints_rest_required_fields(request_type=index_service.Upser assert sorted(expected_params) == sorted(actual_params) -def test_upsert_datapoints_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upsert_datapoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("index", ))) - - def test_remove_datapoints_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -4322,16 +4130,18 @@ def test_remove_datapoints_rest_required_fields(request_type=index_service.Remov # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_datapoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRemoveDatapoints, + "_BaseRemoveDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["index"] = 'index_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_datapoints._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "index" in jsonified_request assert jsonified_request["index"] == 'index_value' @@ -4381,13 +4191,6 @@ def test_remove_datapoints_rest_required_fields(request_type=index_service.Remov assert sorted(expected_params) == sorted(actual_params) -def test_remove_datapoints_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.remove_datapoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("index", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.IndexServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1/test_job_service.py b/tests/unit/gapic/aiplatform_v1/test_job_service.py index 49ea3b19ec..fefa32302f 100644 --- a/tests/unit/gapic/aiplatform_v1/test_job_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_job_service.py @@ -121,6 +121,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -164,150 +176,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert JobServiceClient._get_default_mtls_endpoint(None) is None - assert JobServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert JobServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert JobServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert JobServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert JobServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - JobServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert JobServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert JobServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert JobServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert JobServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - JobServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert JobServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - JobServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert JobServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -321,40 +189,6 @@ def test__get_client_cert_source(): assert JobServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert JobServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(JobServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(JobServiceClient)) -@mock.patch.object(JobServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(JobServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = JobServiceClient._DEFAULT_UNIVERSE - default_endpoint = JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert JobServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert JobServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == JobServiceClient.DEFAULT_MTLS_ENDPOINT - assert JobServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert JobServiceClient._get_api_endpoint(None, None, default_universe, "always") == JobServiceClient.DEFAULT_MTLS_ENDPOINT - assert JobServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == JobServiceClient.DEFAULT_MTLS_ENDPOINT - assert JobServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert JobServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - JobServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert JobServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert JobServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert JobServiceClient._get_universe_domain(None, None) == JobServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - JobServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -739,11 +573,12 @@ def test_job_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -786,11 +621,12 @@ def test_job_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -13945,16 +13781,18 @@ def test_create_custom_job_rest_required_fields(request_type=job_service.CreateC # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateCustomJob, + "_BaseCreateCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14004,13 +13842,6 @@ def test_create_custom_job_rest_required_fields(request_type=job_service.CreateC assert sorted(expected_params) == sorted(actual_params) -def test_create_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "customJob", ))) - - def test_create_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14115,16 +13946,18 @@ def test_get_custom_job_rest_required_fields(request_type=job_service.GetCustomJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetCustomJob, + "_BaseGetCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14173,13 +14006,6 @@ def test_get_custom_job_rest_required_fields(request_type=job_service.GetCustomJ assert sorted(expected_params) == sorted(actual_params) -def test_get_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14282,17 +14108,20 @@ def test_list_custom_jobs_rest_required_fields(request_type=job_service.ListCust # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_custom_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListCustomJobs, + "_BaseListCustomJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_custom_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14342,13 +14171,6 @@ def test_list_custom_jobs_rest_required_fields(request_type=job_service.ListCust assert sorted(expected_params) == sorted(actual_params) -def test_list_custom_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_custom_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_custom_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14520,16 +14342,18 @@ def test_delete_custom_job_rest_required_fields(request_type=job_service.DeleteC # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteCustomJob, + "_BaseDeleteCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14575,13 +14399,6 @@ def test_delete_custom_job_rest_required_fields(request_type=job_service.DeleteC assert sorted(expected_params) == sorted(actual_params) -def test_delete_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14682,16 +14499,18 @@ def test_cancel_custom_job_rest_required_fields(request_type=job_service.CancelC # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelCustomJob, + "_BaseCancelCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14738,13 +14557,6 @@ def test_cancel_custom_job_rest_required_fields(request_type=job_service.CancelC assert sorted(expected_params) == sorted(actual_params) -def test_cancel_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14845,16 +14657,18 @@ def test_create_data_labeling_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDataLabelingJob, + "_BaseCreateDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14904,13 +14718,6 @@ def test_create_data_labeling_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_create_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "dataLabelingJob", ))) - - def test_create_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15015,16 +14822,18 @@ def test_get_data_labeling_job_rest_required_fields(request_type=job_service.Get # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDataLabelingJob, + "_BaseGetDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15073,13 +14882,6 @@ def test_get_data_labeling_job_rest_required_fields(request_type=job_service.Get assert sorted(expected_params) == sorted(actual_params) -def test_get_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15182,17 +14984,20 @@ def test_list_data_labeling_jobs_rest_required_fields(request_type=job_service.L # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_labeling_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDataLabelingJobs, + "_BaseListDataLabelingJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_labeling_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -15242,13 +15047,6 @@ def test_list_data_labeling_jobs_rest_required_fields(request_type=job_service.L assert sorted(expected_params) == sorted(actual_params) -def test_list_data_labeling_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_data_labeling_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_data_labeling_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15420,16 +15218,18 @@ def test_delete_data_labeling_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDataLabelingJob, + "_BaseDeleteDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15475,13 +15275,6 @@ def test_delete_data_labeling_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_delete_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15582,16 +15375,18 @@ def test_cancel_data_labeling_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelDataLabelingJob, + "_BaseCancelDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15638,13 +15433,6 @@ def test_cancel_data_labeling_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_cancel_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15745,16 +15533,18 @@ def test_create_hyperparameter_tuning_job_rest_required_fields(request_type=job_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateHyperparameterTuningJob, + "_BaseCreateHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -15804,13 +15594,6 @@ def test_create_hyperparameter_tuning_job_rest_required_fields(request_type=job_ assert sorted(expected_params) == sorted(actual_params) -def test_create_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "hyperparameterTuningJob", ))) - - def test_create_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15915,16 +15698,18 @@ def test_get_hyperparameter_tuning_job_rest_required_fields(request_type=job_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetHyperparameterTuningJob, + "_BaseGetHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15973,13 +15758,6 @@ def test_get_hyperparameter_tuning_job_rest_required_fields(request_type=job_ser assert sorted(expected_params) == sorted(actual_params) -def test_get_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16082,17 +15860,20 @@ def test_list_hyperparameter_tuning_jobs_rest_required_fields(request_type=job_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_hyperparameter_tuning_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListHyperparameterTuningJobs, + "_BaseListHyperparameterTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_hyperparameter_tuning_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -16142,13 +15923,6 @@ def test_list_hyperparameter_tuning_jobs_rest_required_fields(request_type=job_s assert sorted(expected_params) == sorted(actual_params) -def test_list_hyperparameter_tuning_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_hyperparameter_tuning_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_hyperparameter_tuning_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16320,16 +16094,18 @@ def test_delete_hyperparameter_tuning_job_rest_required_fields(request_type=job_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteHyperparameterTuningJob, + "_BaseDeleteHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16375,13 +16151,6 @@ def test_delete_hyperparameter_tuning_job_rest_required_fields(request_type=job_ assert sorted(expected_params) == sorted(actual_params) -def test_delete_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16482,16 +16251,18 @@ def test_cancel_hyperparameter_tuning_job_rest_required_fields(request_type=job_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelHyperparameterTuningJob, + "_BaseCancelHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16538,13 +16309,6 @@ def test_cancel_hyperparameter_tuning_job_rest_required_fields(request_type=job_ assert sorted(expected_params) == sorted(actual_params) -def test_cancel_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16645,16 +16409,18 @@ def test_create_nas_job_rest_required_fields(request_type=job_service.CreateNasJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateNasJob, + "_BaseCreateNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -16704,13 +16470,6 @@ def test_create_nas_job_rest_required_fields(request_type=job_service.CreateNasJ assert sorted(expected_params) == sorted(actual_params) -def test_create_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "nasJob", ))) - - def test_create_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16815,16 +16574,18 @@ def test_get_nas_job_rest_required_fields(request_type=job_service.GetNasJobRequ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNasJob, + "_BaseGetNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16873,13 +16634,6 @@ def test_get_nas_job_rest_required_fields(request_type=job_service.GetNasJobRequ assert sorted(expected_params) == sorted(actual_params) -def test_get_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16982,17 +16736,20 @@ def test_list_nas_jobs_rest_required_fields(request_type=job_service.ListNasJobs # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNasJobs, + "_BaseListNasJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17042,13 +16799,6 @@ def test_list_nas_jobs_rest_required_fields(request_type=job_service.ListNasJobs assert sorted(expected_params) == sorted(actual_params) -def test_list_nas_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_nas_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_nas_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17220,16 +16970,18 @@ def test_delete_nas_job_rest_required_fields(request_type=job_service.DeleteNasJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNasJob, + "_BaseDeleteNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17275,13 +17027,6 @@ def test_delete_nas_job_rest_required_fields(request_type=job_service.DeleteNasJ assert sorted(expected_params) == sorted(actual_params) -def test_delete_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17382,16 +17127,18 @@ def test_cancel_nas_job_rest_required_fields(request_type=job_service.CancelNasJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelNasJob, + "_BaseCancelNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17438,13 +17185,6 @@ def test_cancel_nas_job_rest_required_fields(request_type=job_service.CancelNasJ assert sorted(expected_params) == sorted(actual_params) -def test_cancel_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17545,16 +17285,18 @@ def test_get_nas_trial_detail_rest_required_fields(request_type=job_service.GetN # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_trial_detail._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNasTrialDetail, + "_BaseGetNasTrialDetail__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_trial_detail._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17603,13 +17345,6 @@ def test_get_nas_trial_detail_rest_required_fields(request_type=job_service.GetN assert sorted(expected_params) == sorted(actual_params) -def test_get_nas_trial_detail_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_nas_trial_detail._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_nas_trial_detail_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17712,17 +17447,20 @@ def test_list_nas_trial_details_rest_required_fields(request_type=job_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_trial_details._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNasTrialDetails, + "_BaseListNasTrialDetails__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_trial_details._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17772,13 +17510,6 @@ def test_list_nas_trial_details_rest_required_fields(request_type=job_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_nas_trial_details_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_nas_trial_details._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_nas_trial_details_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17946,16 +17677,18 @@ def test_create_batch_prediction_job_rest_required_fields(request_type=job_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateBatchPredictionJob, + "_BaseCreateBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -18005,13 +17738,6 @@ def test_create_batch_prediction_job_rest_required_fields(request_type=job_servi assert sorted(expected_params) == sorted(actual_params) -def test_create_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "batchPredictionJob", ))) - - def test_create_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18116,16 +17842,18 @@ def test_get_batch_prediction_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetBatchPredictionJob, + "_BaseGetBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -18174,13 +17902,6 @@ def test_get_batch_prediction_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18283,17 +18004,20 @@ def test_list_batch_prediction_jobs_rest_required_fields(request_type=job_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_batch_prediction_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListBatchPredictionJobs, + "_BaseListBatchPredictionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_batch_prediction_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -18343,13 +18067,6 @@ def test_list_batch_prediction_jobs_rest_required_fields(request_type=job_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_batch_prediction_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_batch_prediction_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_batch_prediction_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18521,16 +18238,18 @@ def test_delete_batch_prediction_job_rest_required_fields(request_type=job_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteBatchPredictionJob, + "_BaseDeleteBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -18576,13 +18295,6 @@ def test_delete_batch_prediction_job_rest_required_fields(request_type=job_servi assert sorted(expected_params) == sorted(actual_params) -def test_delete_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18683,16 +18395,18 @@ def test_cancel_batch_prediction_job_rest_required_fields(request_type=job_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelBatchPredictionJob, + "_BaseCancelBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -18739,13 +18453,6 @@ def test_cancel_batch_prediction_job_rest_required_fields(request_type=job_servi assert sorted(expected_params) == sorted(actual_params) -def test_cancel_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18846,16 +18553,18 @@ def test_create_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateModelDeploymentMonitoringJob, + "_BaseCreateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -18905,13 +18614,6 @@ def test_create_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_create_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "modelDeploymentMonitoringJob", ))) - - def test_create_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19017,7 +18719,12 @@ def test_search_model_deployment_monitoring_stats_anomalies_rest_required_fields # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_deployment_monitoring_stats_anomalies._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchModelDeploymentMonitoringStatsAnomalies, + "_BaseSearchModelDeploymentMonitoringStatsAnomalies__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -19025,9 +18732,6 @@ def test_search_model_deployment_monitoring_stats_anomalies_rest_required_fields jsonified_request["modelDeploymentMonitoringJob"] = 'model_deployment_monitoring_job_value' jsonified_request["deployedModelId"] = 'deployed_model_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_deployment_monitoring_stats_anomalies._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "modelDeploymentMonitoringJob" in jsonified_request assert jsonified_request["modelDeploymentMonitoringJob"] == 'model_deployment_monitoring_job_value' @@ -19079,13 +18783,6 @@ def test_search_model_deployment_monitoring_stats_anomalies_rest_required_fields assert sorted(expected_params) == sorted(actual_params) -def test_search_model_deployment_monitoring_stats_anomalies_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_model_deployment_monitoring_stats_anomalies._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("modelDeploymentMonitoringJob", "deployedModelId", "objectives", ))) - - def test_search_model_deployment_monitoring_stats_anomalies_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19255,16 +18952,18 @@ def test_get_model_deployment_monitoring_job_rest_required_fields(request_type=j # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelDeploymentMonitoringJob, + "_BaseGetModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -19313,13 +19012,6 @@ def test_get_model_deployment_monitoring_job_rest_required_fields(request_type=j assert sorted(expected_params) == sorted(actual_params) -def test_get_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19422,17 +19114,20 @@ def test_list_model_deployment_monitoring_jobs_rest_required_fields(request_type # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_deployment_monitoring_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelDeploymentMonitoringJobs, + "_BaseListModelDeploymentMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_deployment_monitoring_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -19482,13 +19177,6 @@ def test_list_model_deployment_monitoring_jobs_rest_required_fields(request_type assert sorted(expected_params) == sorted(actual_params) -def test_list_model_deployment_monitoring_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_deployment_monitoring_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_deployment_monitoring_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19659,15 +19347,18 @@ def test_update_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateModelDeploymentMonitoringJob, + "_BaseUpdateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -19713,13 +19404,6 @@ def test_update_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_update_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("modelDeploymentMonitoringJob", "updateMask", ))) - - def test_update_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19826,16 +19510,18 @@ def test_delete_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModelDeploymentMonitoringJob, + "_BaseDeleteModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -19881,13 +19567,6 @@ def test_delete_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19988,16 +19667,18 @@ def test_pause_model_deployment_monitoring_job_rest_required_fields(request_type # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePauseModelDeploymentMonitoringJob, + "_BasePauseModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -20044,13 +19725,6 @@ def test_pause_model_deployment_monitoring_job_rest_required_fields(request_type assert sorted(expected_params) == sorted(actual_params) -def test_pause_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.pause_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_pause_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -20151,16 +19825,18 @@ def test_resume_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseResumeModelDeploymentMonitoringJob, + "_BaseResumeModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -20207,13 +19883,6 @@ def test_resume_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_resume_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.resume_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_resume_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_llm_utility_service.py b/tests/unit/gapic/aiplatform_v1/test_llm_utility_service.py index 507230364c..5b5f62e7fe 100644 --- a/tests/unit/gapic/aiplatform_v1/test_llm_utility_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_llm_utility_service.py @@ -84,6 +84,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -127,150 +139,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert LlmUtilityServiceClient._get_default_mtls_endpoint(None) is None - assert LlmUtilityServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert LlmUtilityServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert LlmUtilityServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - LlmUtilityServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert LlmUtilityServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - LlmUtilityServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - LlmUtilityServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -284,40 +152,6 @@ def test__get_client_cert_source(): assert LlmUtilityServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert LlmUtilityServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(LlmUtilityServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(LlmUtilityServiceClient)) -@mock.patch.object(LlmUtilityServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(LlmUtilityServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = LlmUtilityServiceClient._DEFAULT_UNIVERSE - default_endpoint = LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert LlmUtilityServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert LlmUtilityServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - assert LlmUtilityServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert LlmUtilityServiceClient._get_api_endpoint(None, None, default_universe, "always") == LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - assert LlmUtilityServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - assert LlmUtilityServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert LlmUtilityServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - LlmUtilityServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert LlmUtilityServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert LlmUtilityServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert LlmUtilityServiceClient._get_universe_domain(None, None) == LlmUtilityServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - LlmUtilityServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -702,11 +536,12 @@ def test_llm_utility_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -749,11 +584,12 @@ def test_llm_utility_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1676,16 +1512,18 @@ def test_count_tokens_rest_required_fields(request_type=prediction_service.Count # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).count_tokens._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCountTokens, + "_BaseCountTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).count_tokens._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -1735,13 +1573,6 @@ def test_count_tokens_rest_required_fields(request_type=prediction_service.Count assert sorted(expected_params) == sorted(actual_params) -def test_count_tokens_rest_unset_required_fields(): - transport = transports.LlmUtilityServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.count_tokens._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_count_tokens_rest_flattened(): client = LlmUtilityServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -1846,16 +1677,18 @@ def test_compute_tokens_rest_required_fields(request_type=llm_utility_service.Co # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).compute_tokens._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseComputeTokens, + "_BaseComputeTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).compute_tokens._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -1905,13 +1738,6 @@ def test_compute_tokens_rest_required_fields(request_type=llm_utility_service.Co assert sorted(expected_params) == sorted(actual_params) -def test_compute_tokens_rest_unset_required_fields(): - transport = transports.LlmUtilityServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.compute_tokens._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_compute_tokens_rest_flattened(): client = LlmUtilityServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_match_service.py b/tests/unit/gapic/aiplatform_v1/test_match_service.py index 6e3db1f1cb..8e271b5638 100644 --- a/tests/unit/gapic/aiplatform_v1/test_match_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_match_service.py @@ -80,6 +80,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -123,150 +135,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert MatchServiceClient._get_default_mtls_endpoint(None) is None - assert MatchServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert MatchServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert MatchServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MatchServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MatchServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - MatchServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert MatchServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert MatchServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert MatchServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert MatchServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - MatchServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert MatchServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - MatchServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert MatchServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -280,40 +148,6 @@ def test__get_client_cert_source(): assert MatchServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert MatchServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(MatchServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MatchServiceClient)) -@mock.patch.object(MatchServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MatchServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = MatchServiceClient._DEFAULT_UNIVERSE - default_endpoint = MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert MatchServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert MatchServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == MatchServiceClient.DEFAULT_MTLS_ENDPOINT - assert MatchServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert MatchServiceClient._get_api_endpoint(None, None, default_universe, "always") == MatchServiceClient.DEFAULT_MTLS_ENDPOINT - assert MatchServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == MatchServiceClient.DEFAULT_MTLS_ENDPOINT - assert MatchServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert MatchServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - MatchServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert MatchServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert MatchServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert MatchServiceClient._get_universe_domain(None, None) == MatchServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - MatchServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -698,11 +532,12 @@ def test_match_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -745,11 +580,12 @@ def test_match_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1480,16 +1316,18 @@ def test_find_neighbors_rest_required_fields(request_type=match_service.FindNeig # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).find_neighbors._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseFindNeighbors, + "_BaseFindNeighbors__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).find_neighbors._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -1539,13 +1377,6 @@ def test_find_neighbors_rest_required_fields(request_type=match_service.FindNeig assert sorted(expected_params) == sorted(actual_params) -def test_find_neighbors_rest_unset_required_fields(): - transport = transports.MatchServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.find_neighbors._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", ))) - - def test_read_index_datapoints_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -1594,16 +1425,18 @@ def test_read_index_datapoints_rest_required_fields(request_type=match_service.R # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_index_datapoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadIndexDatapoints, + "_BaseReadIndexDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_index_datapoints._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -1653,13 +1486,6 @@ def test_read_index_datapoints_rest_required_fields(request_type=match_service.R assert sorted(expected_params) == sorted(actual_params) -def test_read_index_datapoints_rest_unset_required_fields(): - transport = transports.MatchServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_index_datapoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.MatchServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1/test_metadata_service.py b/tests/unit/gapic/aiplatform_v1/test_metadata_service.py index 1397df7516..3247a3b82a 100644 --- a/tests/unit/gapic/aiplatform_v1/test_metadata_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_metadata_service.py @@ -101,6 +101,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -144,150 +156,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert MetadataServiceClient._get_default_mtls_endpoint(None) is None - assert MetadataServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert MetadataServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MetadataServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - MetadataServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert MetadataServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert MetadataServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert MetadataServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - MetadataServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - MetadataServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert MetadataServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -301,40 +169,6 @@ def test__get_client_cert_source(): assert MetadataServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert MetadataServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(MetadataServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MetadataServiceClient)) -@mock.patch.object(MetadataServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MetadataServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = MetadataServiceClient._DEFAULT_UNIVERSE - default_endpoint = MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert MetadataServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert MetadataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - assert MetadataServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert MetadataServiceClient._get_api_endpoint(None, None, default_universe, "always") == MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - assert MetadataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - assert MetadataServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert MetadataServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - MetadataServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert MetadataServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert MetadataServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert MetadataServiceClient._get_universe_domain(None, None) == MetadataServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - MetadataServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -719,11 +553,12 @@ def test_metadata_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -766,11 +601,12 @@ def test_metadata_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -12460,17 +12296,20 @@ def test_create_metadata_store_rest_required_fields(request_type=metadata_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateMetadataStore, + "_BaseCreateMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("metadata_store_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("metadataStoreId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12518,13 +12357,6 @@ def test_create_metadata_store_rest_required_fields(request_type=metadata_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_metadata_store_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_metadata_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("metadataStoreId", )) & set(("parent", "metadataStore", ))) - - def test_create_metadata_store_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12629,16 +12461,18 @@ def test_get_metadata_store_rest_required_fields(request_type=metadata_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetMetadataStore, + "_BaseGetMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_store._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -12687,13 +12521,6 @@ def test_get_metadata_store_rest_required_fields(request_type=metadata_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_metadata_store_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_metadata_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_metadata_store_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12796,17 +12623,20 @@ def test_list_metadata_stores_rest_required_fields(request_type=metadata_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_stores._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListMetadataStores, + "_BaseListMetadataStores__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_stores._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12856,13 +12686,6 @@ def test_list_metadata_stores_rest_required_fields(request_type=metadata_service assert sorted(expected_params) == sorted(actual_params) -def test_list_metadata_stores_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_metadata_stores._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_metadata_stores_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13034,17 +12857,20 @@ def test_delete_metadata_store_rest_required_fields(request_type=metadata_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_metadata_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteMetadataStore, + "_BaseDeleteMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_metadata_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -13091,13 +12917,6 @@ def test_delete_metadata_store_rest_required_fields(request_type=metadata_servic assert sorted(expected_params) == sorted(actual_params) -def test_delete_metadata_store_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_metadata_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_metadata_store_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13198,17 +13017,20 @@ def test_create_artifact_rest_required_fields(request_type=metadata_service.Crea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateArtifact, + "_BaseCreateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_artifact._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("artifact_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("artifactId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13259,13 +13081,6 @@ def test_create_artifact_rest_required_fields(request_type=metadata_service.Crea assert sorted(expected_params) == sorted(actual_params) -def test_create_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(("artifactId", )) & set(("parent", "artifact", ))) - - def test_create_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13372,16 +13187,18 @@ def test_get_artifact_rest_required_fields(request_type=metadata_service.GetArti # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetArtifact, + "_BaseGetArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_artifact._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13430,13 +13247,6 @@ def test_get_artifact_rest_required_fields(request_type=metadata_service.GetArti assert sorted(expected_params) == sorted(actual_params) -def test_get_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13539,17 +13349,20 @@ def test_list_artifacts_rest_required_fields(request_type=metadata_service.ListA # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_artifacts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListArtifacts, + "_BaseListArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_artifacts._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13599,13 +13412,6 @@ def test_list_artifacts_rest_required_fields(request_type=metadata_service.ListA assert sorted(expected_params) == sorted(actual_params) -def test_list_artifacts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_artifacts._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_artifacts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13772,15 +13578,18 @@ def test_update_artifact_rest_required_fields(request_type=metadata_service.Upda # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateArtifact, + "_BaseUpdateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_artifact._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("allow_missing", "update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("allowMissing", "updateMask", )) # verify required fields with non-default values are left alone @@ -13829,13 +13638,6 @@ def test_update_artifact_rest_required_fields(request_type=metadata_service.Upda assert sorted(expected_params) == sorted(actual_params) -def test_update_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(("allowMissing", "updateMask", )) & set(("artifact", ))) - - def test_update_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13944,17 +13746,20 @@ def test_delete_artifact_rest_required_fields(request_type=metadata_service.Dele # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteArtifact, + "_BaseDeleteArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_artifact._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("etag", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -14001,13 +13806,6 @@ def test_delete_artifact_rest_required_fields(request_type=metadata_service.Dele assert sorted(expected_params) == sorted(actual_params) -def test_delete_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(("etag", )) & set(("name", ))) - - def test_delete_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14113,7 +13911,12 @@ def test_purge_artifacts_rest_required_fields(request_type=metadata_service.Purg # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_artifacts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePurgeArtifacts, + "_BasePurgeArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -14121,9 +13924,6 @@ def test_purge_artifacts_rest_required_fields(request_type=metadata_service.Purg jsonified_request["parent"] = 'parent_value' jsonified_request["filter"] = 'filter_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_artifacts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14172,13 +13972,6 @@ def test_purge_artifacts_rest_required_fields(request_type=metadata_service.Purg assert sorted(expected_params) == sorted(actual_params) -def test_purge_artifacts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.purge_artifacts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "filter", ))) - - def test_purge_artifacts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14279,17 +14072,20 @@ def test_create_context_rest_required_fields(request_type=metadata_service.Creat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateContext, + "_BaseCreateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_context._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("context_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("contextId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14340,13 +14136,6 @@ def test_create_context_rest_required_fields(request_type=metadata_service.Creat assert sorted(expected_params) == sorted(actual_params) -def test_create_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(("contextId", )) & set(("parent", "context", ))) - - def test_create_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14453,16 +14242,18 @@ def test_get_context_rest_required_fields(request_type=metadata_service.GetConte # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetContext, + "_BaseGetContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_context._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14511,13 +14302,6 @@ def test_get_context_rest_required_fields(request_type=metadata_service.GetConte assert sorted(expected_params) == sorted(actual_params) -def test_get_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14620,17 +14404,20 @@ def test_list_contexts_rest_required_fields(request_type=metadata_service.ListCo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListContexts, + "_BaseListContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_contexts._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14680,13 +14467,6 @@ def test_list_contexts_rest_required_fields(request_type=metadata_service.ListCo assert sorted(expected_params) == sorted(actual_params) -def test_list_contexts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_contexts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14853,15 +14633,18 @@ def test_update_context_rest_required_fields(request_type=metadata_service.Updat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateContext, + "_BaseUpdateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_context._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("allow_missing", "update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("allowMissing", "updateMask", )) # verify required fields with non-default values are left alone @@ -14910,13 +14693,6 @@ def test_update_context_rest_required_fields(request_type=metadata_service.Updat assert sorted(expected_params) == sorted(actual_params) -def test_update_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(("allowMissing", "updateMask", )) & set(("context", ))) - - def test_update_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15025,17 +14801,20 @@ def test_delete_context_rest_required_fields(request_type=metadata_service.Delet # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteContext, + "_BaseDeleteContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_context._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("etag", "force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -15082,13 +14861,6 @@ def test_delete_context_rest_required_fields(request_type=metadata_service.Delet assert sorted(expected_params) == sorted(actual_params) -def test_delete_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(("etag", "force", )) & set(("name", ))) - - def test_delete_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15194,7 +14966,12 @@ def test_purge_contexts_rest_required_fields(request_type=metadata_service.Purge # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePurgeContexts, + "_BasePurgeContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -15202,9 +14979,6 @@ def test_purge_contexts_rest_required_fields(request_type=metadata_service.Purge jsonified_request["parent"] = 'parent_value' jsonified_request["filter"] = 'filter_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -15253,13 +15027,6 @@ def test_purge_contexts_rest_required_fields(request_type=metadata_service.Purge assert sorted(expected_params) == sorted(actual_params) -def test_purge_contexts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.purge_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "filter", ))) - - def test_purge_contexts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15360,16 +15127,18 @@ def test_add_context_artifacts_and_executions_rest_required_fields(request_type= # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_artifacts_and_executions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddContextArtifactsAndExecutions, + "_BaseAddContextArtifactsAndExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_artifacts_and_executions._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15419,13 +15188,6 @@ def test_add_context_artifacts_and_executions_rest_required_fields(request_type= assert sorted(expected_params) == sorted(actual_params) -def test_add_context_artifacts_and_executions_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_context_artifacts_and_executions._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_add_context_artifacts_and_executions_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15532,16 +15294,18 @@ def test_add_context_children_rest_required_fields(request_type=metadata_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_children._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddContextChildren, + "_BaseAddContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_children._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15591,13 +15355,6 @@ def test_add_context_children_rest_required_fields(request_type=metadata_service assert sorted(expected_params) == sorted(actual_params) -def test_add_context_children_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_context_children._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_add_context_children_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15702,16 +15459,18 @@ def test_remove_context_children_rest_required_fields(request_type=metadata_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_context_children._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRemoveContextChildren, + "_BaseRemoveContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_context_children._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15761,13 +15520,6 @@ def test_remove_context_children_rest_required_fields(request_type=metadata_serv assert sorted(expected_params) == sorted(actual_params) -def test_remove_context_children_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.remove_context_children._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_remove_context_children_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15872,16 +15624,18 @@ def test_query_context_lineage_subgraph_rest_required_fields(request_type=metada # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_context_lineage_subgraph._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryContextLineageSubgraph, + "_BaseQueryContextLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_context_lineage_subgraph._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15930,13 +15684,6 @@ def test_query_context_lineage_subgraph_rest_required_fields(request_type=metada assert sorted(expected_params) == sorted(actual_params) -def test_query_context_lineage_subgraph_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_context_lineage_subgraph._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_query_context_lineage_subgraph_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16039,17 +15786,20 @@ def test_create_execution_rest_required_fields(request_type=metadata_service.Cre # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateExecution, + "_BaseCreateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_execution._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("execution_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("executionId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -16100,13 +15850,6 @@ def test_create_execution_rest_required_fields(request_type=metadata_service.Cre assert sorted(expected_params) == sorted(actual_params) -def test_create_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(("executionId", )) & set(("parent", "execution", ))) - - def test_create_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16213,16 +15956,18 @@ def test_get_execution_rest_required_fields(request_type=metadata_service.GetExe # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetExecution, + "_BaseGetExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_execution._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16271,13 +16016,6 @@ def test_get_execution_rest_required_fields(request_type=metadata_service.GetExe assert sorted(expected_params) == sorted(actual_params) -def test_get_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16380,17 +16118,20 @@ def test_list_executions_rest_required_fields(request_type=metadata_service.List # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_executions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListExecutions, + "_BaseListExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_executions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -16440,13 +16181,6 @@ def test_list_executions_rest_required_fields(request_type=metadata_service.List assert sorted(expected_params) == sorted(actual_params) -def test_list_executions_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_executions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_executions_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16613,15 +16347,18 @@ def test_update_execution_rest_required_fields(request_type=metadata_service.Upd # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateExecution, + "_BaseUpdateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_execution._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("allow_missing", "update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("allowMissing", "updateMask", )) # verify required fields with non-default values are left alone @@ -16670,13 +16407,6 @@ def test_update_execution_rest_required_fields(request_type=metadata_service.Upd assert sorted(expected_params) == sorted(actual_params) -def test_update_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(("allowMissing", "updateMask", )) & set(("execution", ))) - - def test_update_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16785,17 +16515,20 @@ def test_delete_execution_rest_required_fields(request_type=metadata_service.Del # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteExecution, + "_BaseDeleteExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_execution._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("etag", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -16842,13 +16575,6 @@ def test_delete_execution_rest_required_fields(request_type=metadata_service.Del assert sorted(expected_params) == sorted(actual_params) -def test_delete_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(("etag", )) & set(("name", ))) - - def test_delete_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16954,7 +16680,12 @@ def test_purge_executions_rest_required_fields(request_type=metadata_service.Pur # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_executions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePurgeExecutions, + "_BasePurgeExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -16962,9 +16693,6 @@ def test_purge_executions_rest_required_fields(request_type=metadata_service.Pur jsonified_request["parent"] = 'parent_value' jsonified_request["filter"] = 'filter_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_executions._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -17013,13 +16741,6 @@ def test_purge_executions_rest_required_fields(request_type=metadata_service.Pur assert sorted(expected_params) == sorted(actual_params) -def test_purge_executions_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.purge_executions._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "filter", ))) - - def test_purge_executions_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17120,16 +16841,18 @@ def test_add_execution_events_rest_required_fields(request_type=metadata_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_execution_events._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddExecutionEvents, + "_BaseAddExecutionEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["execution"] = 'execution_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_execution_events._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "execution" in jsonified_request assert jsonified_request["execution"] == 'execution_value' @@ -17179,13 +16902,6 @@ def test_add_execution_events_rest_required_fields(request_type=metadata_service assert sorted(expected_params) == sorted(actual_params) -def test_add_execution_events_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_execution_events._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("execution", ))) - - def test_add_execution_events_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17290,16 +17006,18 @@ def test_query_execution_inputs_and_outputs_rest_required_fields(request_type=me # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_execution_inputs_and_outputs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryExecutionInputsAndOutputs, + "_BaseQueryExecutionInputsAndOutputs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["execution"] = 'execution_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_execution_inputs_and_outputs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "execution" in jsonified_request assert jsonified_request["execution"] == 'execution_value' @@ -17348,13 +17066,6 @@ def test_query_execution_inputs_and_outputs_rest_required_fields(request_type=me assert sorted(expected_params) == sorted(actual_params) -def test_query_execution_inputs_and_outputs_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_execution_inputs_and_outputs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("execution", ))) - - def test_query_execution_inputs_and_outputs_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17457,17 +17168,20 @@ def test_create_metadata_schema_rest_required_fields(request_type=metadata_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_schema._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateMetadataSchema, + "_BaseCreateMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_schema._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("metadata_schema_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("metadataSchemaId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17518,13 +17232,6 @@ def test_create_metadata_schema_rest_required_fields(request_type=metadata_servi assert sorted(expected_params) == sorted(actual_params) -def test_create_metadata_schema_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_metadata_schema._get_unset_required_fields({}) - assert set(unset_fields) == (set(("metadataSchemaId", )) & set(("parent", "metadataSchema", ))) - - def test_create_metadata_schema_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17631,16 +17338,18 @@ def test_get_metadata_schema_rest_required_fields(request_type=metadata_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_schema._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetMetadataSchema, + "_BaseGetMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_schema._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17689,13 +17398,6 @@ def test_get_metadata_schema_rest_required_fields(request_type=metadata_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_metadata_schema_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_metadata_schema._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_metadata_schema_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17798,17 +17500,20 @@ def test_list_metadata_schemas_rest_required_fields(request_type=metadata_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_schemas._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListMetadataSchemas, + "_BaseListMetadataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_schemas._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17858,13 +17563,6 @@ def test_list_metadata_schemas_rest_required_fields(request_type=metadata_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_metadata_schemas_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_metadata_schemas._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_metadata_schemas_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18032,17 +17730,20 @@ def test_query_artifact_lineage_subgraph_rest_required_fields(request_type=metad # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_artifact_lineage_subgraph._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryArtifactLineageSubgraph, + "_BaseQueryArtifactLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["artifact"] = 'artifact_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_artifact_lineage_subgraph._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "max_hops", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "maxHops", )) # verify required fields with non-default values are left alone assert "artifact" in jsonified_request @@ -18092,13 +17793,6 @@ def test_query_artifact_lineage_subgraph_rest_required_fields(request_type=metad assert sorted(expected_params) == sorted(actual_params) -def test_query_artifact_lineage_subgraph_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_artifact_lineage_subgraph._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "maxHops", )) & set(("artifact", ))) - - def test_query_artifact_lineage_subgraph_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_migration_service.py b/tests/unit/gapic/aiplatform_v1/test_migration_service.py index d96f1574d3..850f05bdab 100644 --- a/tests/unit/gapic/aiplatform_v1/test_migration_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_migration_service.py @@ -84,6 +84,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -127,150 +139,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert MigrationServiceClient._get_default_mtls_endpoint(None) is None - assert MigrationServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert MigrationServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MigrationServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - MigrationServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert MigrationServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert MigrationServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert MigrationServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - MigrationServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - MigrationServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert MigrationServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -284,40 +152,6 @@ def test__get_client_cert_source(): assert MigrationServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert MigrationServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(MigrationServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MigrationServiceClient)) -@mock.patch.object(MigrationServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MigrationServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = MigrationServiceClient._DEFAULT_UNIVERSE - default_endpoint = MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert MigrationServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert MigrationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - assert MigrationServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert MigrationServiceClient._get_api_endpoint(None, None, default_universe, "always") == MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - assert MigrationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - assert MigrationServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert MigrationServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - MigrationServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert MigrationServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert MigrationServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert MigrationServiceClient._get_universe_domain(None, None) == MigrationServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - MigrationServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -702,11 +536,12 @@ def test_migration_service_client_get_mtls_endpoint_and_cert_source(client_class for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -749,11 +584,12 @@ def test_migration_service_client_get_mtls_endpoint_and_cert_source(client_class for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1873,16 +1709,18 @@ def test_search_migratable_resources_rest_required_fields(request_type=migration # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_migratable_resources._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchMigratableResources, + "_BaseSearchMigratableResources__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_migratable_resources._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -1932,13 +1770,6 @@ def test_search_migratable_resources_rest_required_fields(request_type=migration assert sorted(expected_params) == sorted(actual_params) -def test_search_migratable_resources_rest_unset_required_fields(): - transport = transports.MigrationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_migratable_resources._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_search_migratable_resources_rest_flattened(): client = MigrationServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2110,16 +1941,18 @@ def test_batch_migrate_resources_rest_required_fields(request_type=migration_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_migrate_resources._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchMigrateResources, + "_BaseBatchMigrateResources__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_migrate_resources._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2166,13 +1999,6 @@ def test_batch_migrate_resources_rest_required_fields(request_type=migration_ser assert sorted(expected_params) == sorted(actual_params) -def test_batch_migrate_resources_rest_unset_required_fields(): - transport = transports.MigrationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_migrate_resources._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "migrateResourceRequests", ))) - - def test_batch_migrate_resources_rest_flattened(): client = MigrationServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_model_garden_service.py b/tests/unit/gapic/aiplatform_v1/test_model_garden_service.py index 44c542bbc4..0d27065d99 100644 --- a/tests/unit/gapic/aiplatform_v1/test_model_garden_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_model_garden_service.py @@ -89,6 +89,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -132,150 +144,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ModelGardenServiceClient._get_default_mtls_endpoint(None) is None - assert ModelGardenServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ModelGardenServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelGardenServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ModelGardenServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ModelGardenServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelGardenServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ModelGardenServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -289,40 +157,6 @@ def test__get_client_cert_source(): assert ModelGardenServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ModelGardenServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ModelGardenServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelGardenServiceClient)) -@mock.patch.object(ModelGardenServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelGardenServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ModelGardenServiceClient._DEFAULT_UNIVERSE - default_endpoint = ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ModelGardenServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ModelGardenServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelGardenServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ModelGardenServiceClient._get_api_endpoint(None, None, default_universe, "always") == ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelGardenServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelGardenServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ModelGardenServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelGardenServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ModelGardenServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ModelGardenServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ModelGardenServiceClient._get_universe_domain(None, None) == ModelGardenServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ModelGardenServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -707,11 +541,12 @@ def test_model_garden_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -754,11 +589,12 @@ def test_model_garden_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1613,17 +1449,20 @@ def test_get_publisher_model_rest_required_fields(request_type=model_garden_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_publisher_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetPublisherModel, + "_BaseGetPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_publisher_model._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("hugging_face_token", "is_hugging_face_model", "language_code", "view", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("huggingFaceToken", "isHuggingFaceModel", "languageCode", "view", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -1673,13 +1512,6 @@ def test_get_publisher_model_rest_required_fields(request_type=model_garden_serv assert sorted(expected_params) == sorted(actual_params) -def test_get_publisher_model_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_publisher_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(("huggingFaceToken", "isHuggingFaceModel", "languageCode", "view", )) & set(("name", ))) - - def test_get_publisher_model_rest_flattened(): client = ModelGardenServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -1786,16 +1618,18 @@ def test_deploy_rest_required_fields(request_type=model_garden_service.DeployReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeploy, + "_BaseDeploy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["destination"] = 'destination_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "destination" in jsonified_request assert jsonified_request["destination"] == 'destination_value' @@ -1842,13 +1676,6 @@ def test_deploy_rest_required_fields(request_type=model_garden_service.DeployReq assert sorted(expected_params) == sorted(actual_params) -def test_deploy_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.deploy._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("destination", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.ModelGardenServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1/test_model_service.py b/tests/unit/gapic/aiplatform_v1/test_model_service.py index 7ebec33bec..d2f7262cca 100644 --- a/tests/unit/gapic/aiplatform_v1/test_model_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_model_service.py @@ -102,6 +102,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -145,150 +157,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ModelServiceClient._get_default_mtls_endpoint(None) is None - assert ModelServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ModelServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ModelServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ModelServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ModelServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ModelServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ModelServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ModelServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ModelServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ModelServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ModelServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -302,40 +170,6 @@ def test__get_client_cert_source(): assert ModelServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ModelServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ModelServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelServiceClient)) -@mock.patch.object(ModelServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ModelServiceClient._DEFAULT_UNIVERSE - default_endpoint = ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ModelServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ModelServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ModelServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ModelServiceClient._get_api_endpoint(None, None, default_universe, "always") == ModelServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ModelServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ModelServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ModelServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ModelServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ModelServiceClient._get_universe_domain(None, None) == ModelServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ModelServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -720,11 +554,12 @@ def test_model_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -767,11 +602,12 @@ def test_model_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -8333,16 +8169,18 @@ def test_upload_model_rest_required_fields(request_type=model_service.UploadMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUploadModel, + "_BaseUploadModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -8389,13 +8227,6 @@ def test_upload_model_rest_required_fields(request_type=model_service.UploadMode assert sorted(expected_params) == sorted(actual_params) -def test_upload_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upload_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "model", ))) - - def test_upload_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8498,16 +8329,18 @@ def test_get_model_rest_required_fields(request_type=model_service.GetModelReque # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModel, + "_BaseGetModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8556,13 +8389,6 @@ def test_get_model_rest_required_fields(request_type=model_service.GetModelReque assert sorted(expected_params) == sorted(actual_params) -def test_get_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8665,17 +8491,20 @@ def test_list_models_rest_required_fields(request_type=model_service.ListModelsR # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_models._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModels, + "_BaseListModels__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_models._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8725,13 +8554,6 @@ def test_list_models_rest_required_fields(request_type=model_service.ListModelsR assert sorted(expected_params) == sorted(actual_params) -def test_list_models_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_models._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_models_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8899,17 +8721,20 @@ def test_list_model_versions_rest_required_fields(request_type=model_service.Lis # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_versions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelVersions, + "_BaseListModelVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_versions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -8959,13 +8784,6 @@ def test_list_model_versions_rest_required_fields(request_type=model_service.Lis assert sorted(expected_params) == sorted(actual_params) -def test_list_model_versions_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_versions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("name", ))) - - def test_list_model_versions_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9133,17 +8951,20 @@ def test_list_model_version_checkpoints_rest_required_fields(request_type=model_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_version_checkpoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelVersionCheckpoints, + "_BaseListModelVersionCheckpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_version_checkpoints._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9193,13 +9014,6 @@ def test_list_model_version_checkpoints_rest_required_fields(request_type=model_ assert sorted(expected_params) == sorted(actual_params) -def test_list_model_version_checkpoints_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_version_checkpoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("name", ))) - - def test_list_model_version_checkpoints_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9366,15 +9180,18 @@ def test_update_model_rest_required_fields(request_type=model_service.UpdateMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateModel, + "_BaseUpdateModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -9423,13 +9240,6 @@ def test_update_model_rest_required_fields(request_type=model_service.UpdateMode assert sorted(expected_params) == sorted(actual_params) -def test_update_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("model", "updateMask", ))) - - def test_update_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9538,16 +9348,18 @@ def test_update_explanation_dataset_rest_required_fields(request_type=model_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_explanation_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateExplanationDataset, + "_BaseUpdateExplanationDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["model"] = 'model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_explanation_dataset._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "model" in jsonified_request assert jsonified_request["model"] == 'model_value' @@ -9594,13 +9406,6 @@ def test_update_explanation_dataset_rest_required_fields(request_type=model_serv assert sorted(expected_params) == sorted(actual_params) -def test_update_explanation_dataset_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_explanation_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("model", ))) - - def test_update_explanation_dataset_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9705,16 +9510,18 @@ def test_delete_model_rest_required_fields(request_type=model_service.DeleteMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModel, + "_BaseDeleteModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9760,13 +9567,6 @@ def test_delete_model_rest_required_fields(request_type=model_service.DeleteMode assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9871,16 +9671,18 @@ def test_delete_model_version_rest_required_fields(request_type=model_service.De # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModelVersion, + "_BaseDeleteModelVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9926,13 +9728,6 @@ def test_delete_model_version_rest_required_fields(request_type=model_service.De assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_version_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_model_version_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10034,7 +9829,12 @@ def test_merge_version_aliases_rest_required_fields(request_type=model_service.M # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).merge_version_aliases._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseMergeVersionAliases, + "_BaseMergeVersionAliases__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -10042,9 +9842,6 @@ def test_merge_version_aliases_rest_required_fields(request_type=model_service.M jsonified_request["name"] = 'name_value' jsonified_request["versionAliases"] = 'version_aliases_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).merge_version_aliases._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10096,13 +9893,6 @@ def test_merge_version_aliases_rest_required_fields(request_type=model_service.M assert sorted(expected_params) == sorted(actual_params) -def test_merge_version_aliases_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.merge_version_aliases._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "versionAliases", ))) - - def test_merge_version_aliases_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10211,16 +10001,18 @@ def test_export_model_rest_required_fields(request_type=model_service.ExportMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportModel, + "_BaseExportModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10267,13 +10059,6 @@ def test_export_model_rest_required_fields(request_type=model_service.ExportMode assert sorted(expected_params) == sorted(actual_params) -def test_export_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "outputConfig", ))) - - def test_export_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10381,7 +10166,12 @@ def test_copy_model_rest_required_fields(request_type=model_service.CopyModelReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).copy_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCopyModel, + "_BaseCopyModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -10389,9 +10179,6 @@ def test_copy_model_rest_required_fields(request_type=model_service.CopyModelReq jsonified_request["parent"] = 'parent_value' jsonified_request["sourceModel"] = 'source_model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).copy_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10440,13 +10227,6 @@ def test_copy_model_rest_required_fields(request_type=model_service.CopyModelReq assert sorted(expected_params) == sorted(actual_params) -def test_copy_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.copy_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "sourceModel", ))) - - def test_copy_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10549,16 +10329,18 @@ def test_import_model_evaluation_rest_required_fields(request_type=model_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_model_evaluation._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportModelEvaluation, + "_BaseImportModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_model_evaluation._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10608,13 +10390,6 @@ def test_import_model_evaluation_rest_required_fields(request_type=model_service assert sorted(expected_params) == sorted(actual_params) -def test_import_model_evaluation_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_model_evaluation._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "modelEvaluation", ))) - - def test_import_model_evaluation_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10719,16 +10494,18 @@ def test_batch_import_model_evaluation_slices_rest_required_fields(request_type= # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_model_evaluation_slices._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchImportModelEvaluationSlices, + "_BaseBatchImportModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_model_evaluation_slices._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10778,13 +10555,6 @@ def test_batch_import_model_evaluation_slices_rest_required_fields(request_type= assert sorted(expected_params) == sorted(actual_params) -def test_batch_import_model_evaluation_slices_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_import_model_evaluation_slices._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "modelEvaluationSlices", ))) - - def test_batch_import_model_evaluation_slices_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10889,16 +10659,18 @@ def test_batch_import_evaluated_annotations_rest_required_fields(request_type=mo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_evaluated_annotations._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchImportEvaluatedAnnotations, + "_BaseBatchImportEvaluatedAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_evaluated_annotations._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10948,13 +10720,6 @@ def test_batch_import_evaluated_annotations_rest_required_fields(request_type=mo assert sorted(expected_params) == sorted(actual_params) -def test_batch_import_evaluated_annotations_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_import_evaluated_annotations._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "evaluatedAnnotations", ))) - - def test_batch_import_evaluated_annotations_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11059,16 +10824,18 @@ def test_get_model_evaluation_rest_required_fields(request_type=model_service.Ge # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelEvaluation, + "_BaseGetModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11117,13 +10884,6 @@ def test_get_model_evaluation_rest_required_fields(request_type=model_service.Ge assert sorted(expected_params) == sorted(actual_params) -def test_get_model_evaluation_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_evaluation._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_evaluation_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11226,17 +10986,20 @@ def test_list_model_evaluations_rest_required_fields(request_type=model_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluations._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelEvaluations, + "_BaseListModelEvaluations__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluations._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11286,13 +11049,6 @@ def test_list_model_evaluations_rest_required_fields(request_type=model_service. assert sorted(expected_params) == sorted(actual_params) -def test_list_model_evaluations_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_evaluations._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_evaluations_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11460,16 +11216,18 @@ def test_get_model_evaluation_slice_rest_required_fields(request_type=model_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation_slice._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelEvaluationSlice, + "_BaseGetModelEvaluationSlice__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation_slice._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11518,13 +11276,6 @@ def test_get_model_evaluation_slice_rest_required_fields(request_type=model_serv assert sorted(expected_params) == sorted(actual_params) -def test_get_model_evaluation_slice_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_evaluation_slice._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_evaluation_slice_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11627,17 +11378,20 @@ def test_list_model_evaluation_slices_rest_required_fields(request_type=model_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluation_slices._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelEvaluationSlices, + "_BaseListModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluation_slices._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11687,13 +11441,6 @@ def test_list_model_evaluation_slices_rest_required_fields(request_type=model_se assert sorted(expected_params) == sorted(actual_params) -def test_list_model_evaluation_slices_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_evaluation_slices._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_evaluation_slices_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_notebook_service.py b/tests/unit/gapic/aiplatform_v1/test_notebook_service.py index 88a95661c3..5cb539cc9f 100644 --- a/tests/unit/gapic/aiplatform_v1/test_notebook_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_notebook_service.py @@ -105,6 +105,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -148,150 +160,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert NotebookServiceClient._get_default_mtls_endpoint(None) is None - assert NotebookServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert NotebookServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert NotebookServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - NotebookServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert NotebookServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert NotebookServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert NotebookServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - NotebookServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - NotebookServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert NotebookServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -305,40 +173,6 @@ def test__get_client_cert_source(): assert NotebookServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert NotebookServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(NotebookServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(NotebookServiceClient)) -@mock.patch.object(NotebookServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(NotebookServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = NotebookServiceClient._DEFAULT_UNIVERSE - default_endpoint = NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert NotebookServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert NotebookServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - assert NotebookServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert NotebookServiceClient._get_api_endpoint(None, None, default_universe, "always") == NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - assert NotebookServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - assert NotebookServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert NotebookServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - NotebookServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert NotebookServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert NotebookServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert NotebookServiceClient._get_universe_domain(None, None) == NotebookServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - NotebookServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -723,11 +557,12 @@ def test_notebook_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -770,11 +605,12 @@ def test_notebook_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -6903,17 +6739,20 @@ def test_create_notebook_runtime_template_rest_required_fields(request_type=note # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateNotebookRuntimeTemplate, + "_BaseCreateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_runtime_template._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("notebook_runtime_template_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("notebookRuntimeTemplateId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6961,13 +6800,6 @@ def test_create_notebook_runtime_template_rest_required_fields(request_type=note assert sorted(expected_params) == sorted(actual_params) -def test_create_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(("notebookRuntimeTemplateId", )) & set(("parent", "notebookRuntimeTemplate", ))) - - def test_create_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7072,16 +6904,18 @@ def test_get_notebook_runtime_template_rest_required_fields(request_type=noteboo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNotebookRuntimeTemplate, + "_BaseGetNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime_template._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7130,13 +6964,6 @@ def test_get_notebook_runtime_template_rest_required_fields(request_type=noteboo assert sorted(expected_params) == sorted(actual_params) -def test_get_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7239,17 +7066,20 @@ def test_list_notebook_runtime_templates_rest_required_fields(request_type=noteb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtime_templates._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNotebookRuntimeTemplates, + "_BaseListNotebookRuntimeTemplates__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtime_templates._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -7299,13 +7129,6 @@ def test_list_notebook_runtime_templates_rest_required_fields(request_type=noteb assert sorted(expected_params) == sorted(actual_params) -def test_list_notebook_runtime_templates_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_notebook_runtime_templates._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_notebook_runtime_templates_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7477,16 +7300,18 @@ def test_delete_notebook_runtime_template_rest_required_fields(request_type=note # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNotebookRuntimeTemplate, + "_BaseDeleteNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime_template._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7532,13 +7357,6 @@ def test_delete_notebook_runtime_template_rest_required_fields(request_type=note assert sorted(expected_params) == sorted(actual_params) -def test_delete_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7638,15 +7456,18 @@ def test_update_notebook_runtime_template_rest_required_fields(request_type=note # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateNotebookRuntimeTemplate, + "_BaseUpdateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_notebook_runtime_template._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -7695,13 +7516,6 @@ def test_update_notebook_runtime_template_rest_required_fields(request_type=note assert sorted(expected_params) == sorted(actual_params) -def test_update_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("notebookRuntimeTemplate", "updateMask", ))) - - def test_update_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7811,7 +7625,12 @@ def test_assign_notebook_runtime_rest_required_fields(request_type=notebook_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assign_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAssignNotebookRuntime, + "_BaseAssignNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -7819,9 +7638,6 @@ def test_assign_notebook_runtime_rest_required_fields(request_type=notebook_serv jsonified_request["parent"] = 'parent_value' jsonified_request["notebookRuntimeTemplate"] = 'notebook_runtime_template_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assign_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -7870,13 +7686,6 @@ def test_assign_notebook_runtime_rest_required_fields(request_type=notebook_serv assert sorted(expected_params) == sorted(actual_params) -def test_assign_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.assign_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "notebookRuntimeTemplate", "notebookRuntime", ))) - - def test_assign_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7983,16 +7792,18 @@ def test_get_notebook_runtime_rest_required_fields(request_type=notebook_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNotebookRuntime, + "_BaseGetNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8041,13 +7852,6 @@ def test_get_notebook_runtime_rest_required_fields(request_type=notebook_service assert sorted(expected_params) == sorted(actual_params) -def test_get_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8150,17 +7954,20 @@ def test_list_notebook_runtimes_rest_required_fields(request_type=notebook_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtimes._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNotebookRuntimes, + "_BaseListNotebookRuntimes__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtimes._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8210,13 +8017,6 @@ def test_list_notebook_runtimes_rest_required_fields(request_type=notebook_servi assert sorted(expected_params) == sorted(actual_params) -def test_list_notebook_runtimes_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_notebook_runtimes._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_notebook_runtimes_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8388,16 +8188,18 @@ def test_delete_notebook_runtime_rest_required_fields(request_type=notebook_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNotebookRuntime, + "_BaseDeleteNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8443,13 +8245,6 @@ def test_delete_notebook_runtime_rest_required_fields(request_type=notebook_serv assert sorted(expected_params) == sorted(actual_params) -def test_delete_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8554,16 +8349,18 @@ def test_upgrade_notebook_runtime_rest_required_fields(request_type=notebook_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upgrade_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpgradeNotebookRuntime, + "_BaseUpgradeNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upgrade_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8610,13 +8407,6 @@ def test_upgrade_notebook_runtime_rest_required_fields(request_type=notebook_ser assert sorted(expected_params) == sorted(actual_params) -def test_upgrade_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upgrade_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_upgrade_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8721,16 +8511,18 @@ def test_start_notebook_runtime_rest_required_fields(request_type=notebook_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).start_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStartNotebookRuntime, + "_BaseStartNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).start_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8777,13 +8569,6 @@ def test_start_notebook_runtime_rest_required_fields(request_type=notebook_servi assert sorted(expected_params) == sorted(actual_params) -def test_start_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.start_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_start_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8888,16 +8673,18 @@ def test_stop_notebook_runtime_rest_required_fields(request_type=notebook_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStopNotebookRuntime, + "_BaseStopNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8944,13 +8731,6 @@ def test_stop_notebook_runtime_rest_required_fields(request_type=notebook_servic assert sorted(expected_params) == sorted(actual_params) -def test_stop_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stop_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_stop_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9055,17 +8835,20 @@ def test_create_notebook_execution_job_rest_required_fields(request_type=noteboo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_execution_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateNotebookExecutionJob, + "_BaseCreateNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_execution_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("notebook_execution_job_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("notebookExecutionJobId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9113,13 +8896,6 @@ def test_create_notebook_execution_job_rest_required_fields(request_type=noteboo assert sorted(expected_params) == sorted(actual_params) -def test_create_notebook_execution_job_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_notebook_execution_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("notebookExecutionJobId", )) & set(("parent", "notebookExecutionJob", ))) - - def test_create_notebook_execution_job_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9224,17 +9000,20 @@ def test_get_notebook_execution_job_rest_required_fields(request_type=notebook_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_execution_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNotebookExecutionJob, + "_BaseGetNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_execution_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("view", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9284,13 +9063,6 @@ def test_get_notebook_execution_job_rest_required_fields(request_type=notebook_s assert sorted(expected_params) == sorted(actual_params) -def test_get_notebook_execution_job_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_notebook_execution_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("view", )) & set(("name", ))) - - def test_get_notebook_execution_job_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9393,17 +9165,20 @@ def test_list_notebook_execution_jobs_rest_required_fields(request_type=notebook # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_execution_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNotebookExecutionJobs, + "_BaseListNotebookExecutionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_execution_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "view", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "view", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9453,13 +9228,6 @@ def test_list_notebook_execution_jobs_rest_required_fields(request_type=notebook assert sorted(expected_params) == sorted(actual_params) -def test_list_notebook_execution_jobs_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_notebook_execution_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "view", )) & set(("parent", ))) - - def test_list_notebook_execution_jobs_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9631,16 +9399,18 @@ def test_delete_notebook_execution_job_rest_required_fields(request_type=noteboo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_execution_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNotebookExecutionJob, + "_BaseDeleteNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_execution_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9686,13 +9456,6 @@ def test_delete_notebook_execution_job_rest_required_fields(request_type=noteboo assert sorted(expected_params) == sorted(actual_params) -def test_delete_notebook_execution_job_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_notebook_execution_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_notebook_execution_job_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_persistent_resource_service.py b/tests/unit/gapic/aiplatform_v1/test_persistent_resource_service.py index 1ea2f6dc51..b816e7d055 100644 --- a/tests/unit/gapic/aiplatform_v1/test_persistent_resource_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_persistent_resource_service.py @@ -96,6 +96,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -139,150 +151,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert PersistentResourceServiceClient._get_default_mtls_endpoint(None) is None - assert PersistentResourceServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert PersistentResourceServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PersistentResourceServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - PersistentResourceServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert PersistentResourceServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - PersistentResourceServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - PersistentResourceServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -296,40 +164,6 @@ def test__get_client_cert_source(): assert PersistentResourceServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert PersistentResourceServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(PersistentResourceServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PersistentResourceServiceClient)) -@mock.patch.object(PersistentResourceServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PersistentResourceServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = PersistentResourceServiceClient._DEFAULT_UNIVERSE - default_endpoint = PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert PersistentResourceServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert PersistentResourceServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - assert PersistentResourceServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert PersistentResourceServiceClient._get_api_endpoint(None, None, default_universe, "always") == PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - assert PersistentResourceServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - assert PersistentResourceServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert PersistentResourceServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - PersistentResourceServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert PersistentResourceServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert PersistentResourceServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert PersistentResourceServiceClient._get_universe_domain(None, None) == PersistentResourceServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - PersistentResourceServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -714,11 +548,12 @@ def test_persistent_resource_service_client_get_mtls_endpoint_and_cert_source(cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -761,11 +596,12 @@ def test_persistent_resource_service_client_get_mtls_endpoint_and_cert_source(cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3185,7 +3021,12 @@ def test_create_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped assert "persistentResourceId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreatePersistentResource, + "_BaseCreatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -3195,10 +3036,8 @@ def test_create_persistent_resource_rest_required_fields(request_type=persistent jsonified_request["parent"] = 'parent_value' jsonified_request["persistentResourceId"] = 'persistent_resource_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_persistent_resource._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("persistent_resource_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("persistentResourceId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3252,13 +3091,6 @@ def test_create_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_create_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(("persistentResourceId", )) & set(("parent", "persistentResource", "persistentResourceId", ))) - - def test_create_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3363,16 +3195,18 @@ def test_get_persistent_resource_rest_required_fields(request_type=persistent_re # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetPersistentResource, + "_BaseGetPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_persistent_resource._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3421,13 +3255,6 @@ def test_get_persistent_resource_rest_required_fields(request_type=persistent_re assert sorted(expected_params) == sorted(actual_params) -def test_get_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3530,17 +3357,20 @@ def test_list_persistent_resources_rest_required_fields(request_type=persistent_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_persistent_resources._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListPersistentResources, + "_BaseListPersistentResources__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_persistent_resources._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3590,13 +3420,6 @@ def test_list_persistent_resources_rest_required_fields(request_type=persistent_ assert sorted(expected_params) == sorted(actual_params) -def test_list_persistent_resources_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_persistent_resources._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_persistent_resources_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3768,16 +3591,18 @@ def test_delete_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeletePersistentResource, + "_BaseDeletePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_persistent_resource._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3823,13 +3648,6 @@ def test_delete_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_delete_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3933,15 +3751,18 @@ def test_update_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdatePersistentResource, + "_BaseUpdatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_persistent_resource._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3987,13 +3808,6 @@ def test_update_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_update_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("persistentResource", "updateMask", ))) - - def test_update_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4100,16 +3914,18 @@ def test_reboot_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).reboot_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRebootPersistentResource, + "_BaseRebootPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).reboot_persistent_resource._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4156,13 +3972,6 @@ def test_reboot_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_reboot_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.reboot_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_reboot_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_pipeline_service.py b/tests/unit/gapic/aiplatform_v1/test_pipeline_service.py index 7402191cfd..6774861bfb 100644 --- a/tests/unit/gapic/aiplatform_v1/test_pipeline_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_pipeline_service.py @@ -109,6 +109,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -152,150 +164,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert PipelineServiceClient._get_default_mtls_endpoint(None) is None - assert PipelineServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert PipelineServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PipelineServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - PipelineServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert PipelineServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert PipelineServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert PipelineServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - PipelineServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - PipelineServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert PipelineServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -309,40 +177,6 @@ def test__get_client_cert_source(): assert PipelineServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert PipelineServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(PipelineServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PipelineServiceClient)) -@mock.patch.object(PipelineServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PipelineServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = PipelineServiceClient._DEFAULT_UNIVERSE - default_endpoint = PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert PipelineServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert PipelineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - assert PipelineServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert PipelineServiceClient._get_api_endpoint(None, None, default_universe, "always") == PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - assert PipelineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - assert PipelineServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert PipelineServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - PipelineServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert PipelineServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert PipelineServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert PipelineServiceClient._get_universe_domain(None, None) == PipelineServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - PipelineServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -727,11 +561,12 @@ def test_pipeline_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -774,11 +609,12 @@ def test_pipeline_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5351,16 +5187,18 @@ def test_create_training_pipeline_rest_required_fields(request_type=pipeline_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTrainingPipeline, + "_BaseCreateTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -5410,13 +5248,6 @@ def test_create_training_pipeline_rest_required_fields(request_type=pipeline_ser assert sorted(expected_params) == sorted(actual_params) -def test_create_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "trainingPipeline", ))) - - def test_create_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5521,16 +5352,18 @@ def test_get_training_pipeline_rest_required_fields(request_type=pipeline_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTrainingPipeline, + "_BaseGetTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5579,13 +5412,6 @@ def test_get_training_pipeline_rest_required_fields(request_type=pipeline_servic assert sorted(expected_params) == sorted(actual_params) -def test_get_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5688,17 +5514,20 @@ def test_list_training_pipelines_rest_required_fields(request_type=pipeline_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_training_pipelines._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTrainingPipelines, + "_BaseListTrainingPipelines__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_training_pipelines._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5748,13 +5577,6 @@ def test_list_training_pipelines_rest_required_fields(request_type=pipeline_serv assert sorted(expected_params) == sorted(actual_params) -def test_list_training_pipelines_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_training_pipelines._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_training_pipelines_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5926,16 +5748,18 @@ def test_delete_training_pipeline_rest_required_fields(request_type=pipeline_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTrainingPipeline, + "_BaseDeleteTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5981,13 +5805,6 @@ def test_delete_training_pipeline_rest_required_fields(request_type=pipeline_ser assert sorted(expected_params) == sorted(actual_params) -def test_delete_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6088,16 +5905,18 @@ def test_cancel_training_pipeline_rest_required_fields(request_type=pipeline_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelTrainingPipeline, + "_BaseCancelTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6144,13 +5963,6 @@ def test_cancel_training_pipeline_rest_required_fields(request_type=pipeline_ser assert sorted(expected_params) == sorted(actual_params) -def test_cancel_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6251,17 +6063,20 @@ def test_create_pipeline_job_rest_required_fields(request_type=pipeline_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreatePipelineJob, + "_BaseCreatePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_pipeline_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("pipeline_job_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pipelineJobId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6312,13 +6127,6 @@ def test_create_pipeline_job_rest_required_fields(request_type=pipeline_service. assert sorted(expected_params) == sorted(actual_params) -def test_create_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pipelineJobId", )) & set(("parent", "pipelineJob", ))) - - def test_create_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6425,16 +6233,18 @@ def test_get_pipeline_job_rest_required_fields(request_type=pipeline_service.Get # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetPipelineJob, + "_BaseGetPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_pipeline_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6483,13 +6293,6 @@ def test_get_pipeline_job_rest_required_fields(request_type=pipeline_service.Get assert sorted(expected_params) == sorted(actual_params) -def test_get_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6592,17 +6395,20 @@ def test_list_pipeline_jobs_rest_required_fields(request_type=pipeline_service.L # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_pipeline_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListPipelineJobs, + "_BaseListPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_pipeline_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6652,13 +6458,6 @@ def test_list_pipeline_jobs_rest_required_fields(request_type=pipeline_service.L assert sorted(expected_params) == sorted(actual_params) -def test_list_pipeline_jobs_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_pipeline_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_pipeline_jobs_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6830,16 +6629,18 @@ def test_delete_pipeline_job_rest_required_fields(request_type=pipeline_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeletePipelineJob, + "_BaseDeletePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_pipeline_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6885,13 +6686,6 @@ def test_delete_pipeline_job_rest_required_fields(request_type=pipeline_service. assert sorted(expected_params) == sorted(actual_params) -def test_delete_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6997,7 +6791,12 @@ def test_batch_delete_pipeline_jobs_rest_required_fields(request_type=pipeline_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_pipeline_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchDeletePipelineJobs, + "_BaseBatchDeletePipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -7005,9 +6804,6 @@ def test_batch_delete_pipeline_jobs_rest_required_fields(request_type=pipeline_s jsonified_request["parent"] = 'parent_value' jsonified_request["names"] = 'names_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_pipeline_jobs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -7056,13 +6852,6 @@ def test_batch_delete_pipeline_jobs_rest_required_fields(request_type=pipeline_s assert sorted(expected_params) == sorted(actual_params) -def test_batch_delete_pipeline_jobs_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_delete_pipeline_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "names", ))) - - def test_batch_delete_pipeline_jobs_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7165,16 +6954,18 @@ def test_cancel_pipeline_job_rest_required_fields(request_type=pipeline_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelPipelineJob, + "_BaseCancelPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_pipeline_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7221,13 +7012,6 @@ def test_cancel_pipeline_job_rest_required_fields(request_type=pipeline_service. assert sorted(expected_params) == sorted(actual_params) -def test_cancel_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7333,7 +7117,12 @@ def test_batch_cancel_pipeline_jobs_rest_required_fields(request_type=pipeline_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_cancel_pipeline_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCancelPipelineJobs, + "_BaseBatchCancelPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -7341,9 +7130,6 @@ def test_batch_cancel_pipeline_jobs_rest_required_fields(request_type=pipeline_s jsonified_request["parent"] = 'parent_value' jsonified_request["names"] = 'names_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_cancel_pipeline_jobs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -7392,13 +7178,6 @@ def test_batch_cancel_pipeline_jobs_rest_required_fields(request_type=pipeline_s assert sorted(expected_params) == sorted(actual_params) -def test_batch_cancel_pipeline_jobs_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_cancel_pipeline_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "names", ))) - - def test_batch_cancel_pipeline_jobs_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_prediction_service.py b/tests/unit/gapic/aiplatform_v1/test_prediction_service.py index afa8523211..6e7966adb3 100644 --- a/tests/unit/gapic/aiplatform_v1/test_prediction_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_prediction_service.py @@ -92,6 +92,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -135,150 +147,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert PredictionServiceClient._get_default_mtls_endpoint(None) is None - assert PredictionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert PredictionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PredictionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - PredictionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert PredictionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert PredictionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert PredictionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - PredictionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - PredictionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert PredictionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -292,40 +160,6 @@ def test__get_client_cert_source(): assert PredictionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert PredictionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(PredictionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PredictionServiceClient)) -@mock.patch.object(PredictionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PredictionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = PredictionServiceClient._DEFAULT_UNIVERSE - default_endpoint = PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert PredictionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert PredictionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - assert PredictionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert PredictionServiceClient._get_api_endpoint(None, None, default_universe, "always") == PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - assert PredictionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - assert PredictionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert PredictionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - PredictionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert PredictionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert PredictionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert PredictionServiceClient._get_universe_domain(None, None) == PredictionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - PredictionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -710,11 +544,12 @@ def test_prediction_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -757,11 +592,12 @@ def test_prediction_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -4391,16 +4227,18 @@ def test_predict_rest_required_fields(request_type=prediction_service.PredictReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePredict, + "_BasePredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -4450,13 +4288,6 @@ def test_predict_rest_required_fields(request_type=prediction_service.PredictReq assert sorted(expected_params) == sorted(actual_params) -def test_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "instances", ))) - - def test_predict_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4563,16 +4394,18 @@ def test_raw_predict_rest_required_fields(request_type=prediction_service.RawPre # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).raw_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRawPredict, + "_BaseRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).raw_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -4620,13 +4453,6 @@ def test_raw_predict_rest_required_fields(request_type=prediction_service.RawPre assert sorted(expected_params) == sorted(actual_params) -def test_raw_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.raw_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_raw_predict_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4729,16 +4555,18 @@ def test_stream_raw_predict_rest_required_fields(request_type=prediction_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_raw_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamRawPredict, + "_BaseStreamRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_raw_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -4789,13 +4617,6 @@ def test_stream_raw_predict_rest_required_fields(request_type=prediction_service assert sorted(expected_params) == sorted(actual_params) -def test_stream_raw_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stream_raw_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_stream_raw_predict_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4901,16 +4722,18 @@ def test_direct_predict_rest_required_fields(request_type=prediction_service.Dir # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDirectPredict, + "_BaseDirectPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -4960,13 +4783,6 @@ def test_direct_predict_rest_required_fields(request_type=prediction_service.Dir assert sorted(expected_params) == sorted(actual_params) -def test_direct_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.direct_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_direct_raw_predict_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -5015,16 +4831,18 @@ def test_direct_raw_predict_rest_required_fields(request_type=prediction_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_raw_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDirectRawPredict, + "_BaseDirectRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_raw_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5074,13 +4892,6 @@ def test_direct_raw_predict_rest_required_fields(request_type=prediction_service assert sorted(expected_params) == sorted(actual_params) -def test_direct_raw_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.direct_raw_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_stream_direct_predict_rest_no_http_options(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5162,16 +4973,18 @@ def test_server_streaming_predict_rest_required_fields(request_type=prediction_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).server_streaming_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseServerStreamingPredict, + "_BaseServerStreamingPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).server_streaming_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5224,13 +5037,6 @@ def test_server_streaming_predict_rest_required_fields(request_type=prediction_s assert sorted(expected_params) == sorted(actual_params) -def test_server_streaming_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.server_streaming_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_streaming_raw_predict_rest_no_http_options(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5290,16 +5096,18 @@ def test_explain_rest_required_fields(request_type=prediction_service.ExplainReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).explain._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExplain, + "_BaseExplain__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).explain._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5349,13 +5157,6 @@ def test_explain_rest_required_fields(request_type=prediction_service.ExplainReq assert sorted(expected_params) == sorted(actual_params) -def test_explain_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.explain._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "instances", ))) - - def test_explain_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5464,16 +5265,18 @@ def test_generate_content_rest_required_fields(request_type=prediction_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGenerateContent, + "_BaseGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["model"] = 'model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "model" in jsonified_request assert jsonified_request["model"] == 'model_value' @@ -5523,13 +5326,6 @@ def test_generate_content_rest_required_fields(request_type=prediction_service.G assert sorted(expected_params) == sorted(actual_params) -def test_generate_content_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.generate_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("model", "contents", ))) - - def test_generate_content_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5634,16 +5430,18 @@ def test_stream_generate_content_rest_required_fields(request_type=prediction_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_generate_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamGenerateContent, + "_BaseStreamGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["model"] = 'model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_generate_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "model" in jsonified_request assert jsonified_request["model"] == 'model_value' @@ -5696,13 +5494,6 @@ def test_stream_generate_content_rest_required_fields(request_type=prediction_se assert sorted(expected_params) == sorted(actual_params) -def test_stream_generate_content_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stream_generate_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("model", "contents", ))) - - def test_stream_generate_content_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_execution_service.py b/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_execution_service.py index 9e9c5ecd08..080e9f8a70 100644 --- a/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_execution_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_execution_service.py @@ -85,6 +85,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -128,150 +140,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(None) is None - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ReasoningEngineExecutionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineExecutionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ReasoningEngineExecutionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -285,40 +153,6 @@ def test__get_client_cert_source(): assert ReasoningEngineExecutionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ReasoningEngineExecutionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ReasoningEngineExecutionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineExecutionServiceClient)) -@mock.patch.object(ReasoningEngineExecutionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineExecutionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - default_endpoint = ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, default_universe, "always") == ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ReasoningEngineExecutionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ReasoningEngineExecutionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ReasoningEngineExecutionServiceClient._get_universe_domain(None, None) == ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ReasoningEngineExecutionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -703,11 +537,12 @@ def test_reasoning_engine_execution_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -750,11 +585,12 @@ def test_reasoning_engine_execution_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1949,16 +1785,18 @@ def test_query_reasoning_engine_rest_required_fields(request_type=reasoning_engi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryReasoningEngine, + "_BaseQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2008,13 +1846,6 @@ def test_query_reasoning_engine_rest_required_fields(request_type=reasoning_engi assert sorted(expected_params) == sorted(actual_params) -def test_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_stream_query_reasoning_engine_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -2063,16 +1894,18 @@ def test_stream_query_reasoning_engine_rest_required_fields(request_type=reasoni # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamQueryReasoningEngine, + "_BaseStreamQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2123,13 +1956,6 @@ def test_stream_query_reasoning_engine_rest_required_fields(request_type=reasoni assert sorted(expected_params) == sorted(actual_params) -def test_stream_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stream_query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_async_query_reasoning_engine_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -2182,16 +2008,18 @@ def test_async_query_reasoning_engine_rest_required_fields(request_type=reasonin # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAsyncQueryReasoningEngine, + "_BaseAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2238,13 +2066,6 @@ def test_async_query_reasoning_engine_rest_required_fields(request_type=reasonin assert sorted(expected_params) == sorted(actual_params) -def test_async_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.async_query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_async_query_reasoning_engine_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -2294,7 +2115,12 @@ def test_cancel_async_query_reasoning_engine_rest_required_fields(request_type=r # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_async_query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelAsyncQueryReasoningEngine, + "_BaseCancelAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -2302,9 +2128,6 @@ def test_cancel_async_query_reasoning_engine_rest_required_fields(request_type=r jsonified_request["name"] = 'name_value' jsonified_request["operationName"] = 'operation_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_async_query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2356,13 +2179,6 @@ def test_cancel_async_query_reasoning_engine_rest_required_fields(request_type=r assert sorted(expected_params) == sorted(actual_params) -def test_cancel_async_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_async_query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "operationName", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.ReasoningEngineExecutionServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_service.py b/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_service.py index 96d70711db..33367c7cd8 100644 --- a/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_reasoning_engine_service.py @@ -93,6 +93,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -136,150 +148,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(None) is None - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ReasoningEngineServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ReasoningEngineServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ReasoningEngineServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -293,40 +161,6 @@ def test__get_client_cert_source(): assert ReasoningEngineServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ReasoningEngineServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ReasoningEngineServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineServiceClient)) -@mock.patch.object(ReasoningEngineServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ReasoningEngineServiceClient._DEFAULT_UNIVERSE - default_endpoint = ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ReasoningEngineServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ReasoningEngineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, default_universe, "always") == ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ReasoningEngineServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ReasoningEngineServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ReasoningEngineServiceClient._get_universe_domain(None, None) == ReasoningEngineServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ReasoningEngineServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -711,11 +545,12 @@ def test_reasoning_engine_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -758,11 +593,12 @@ def test_reasoning_engine_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2849,16 +2685,18 @@ def test_create_reasoning_engine_rest_required_fields(request_type=reasoning_eng # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateReasoningEngine, + "_BaseCreateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2905,13 +2743,6 @@ def test_create_reasoning_engine_rest_required_fields(request_type=reasoning_eng assert sorted(expected_params) == sorted(actual_params) -def test_create_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "reasoningEngine", ))) - - def test_create_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3014,16 +2845,18 @@ def test_get_reasoning_engine_rest_required_fields(request_type=reasoning_engine # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetReasoningEngine, + "_BaseGetReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3072,13 +2905,6 @@ def test_get_reasoning_engine_rest_required_fields(request_type=reasoning_engine assert sorted(expected_params) == sorted(actual_params) -def test_get_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3181,17 +3007,20 @@ def test_list_reasoning_engines_rest_required_fields(request_type=reasoning_engi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_reasoning_engines._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListReasoningEngines, + "_BaseListReasoningEngines__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_reasoning_engines._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3241,13 +3070,6 @@ def test_list_reasoning_engines_rest_required_fields(request_type=reasoning_engi assert sorted(expected_params) == sorted(actual_params) -def test_list_reasoning_engines_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_reasoning_engines._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_reasoning_engines_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3418,15 +3240,18 @@ def test_update_reasoning_engine_rest_required_fields(request_type=reasoning_eng # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateReasoningEngine, + "_BaseUpdateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_reasoning_engine._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3472,13 +3297,6 @@ def test_update_reasoning_engine_rest_required_fields(request_type=reasoning_eng assert sorted(expected_params) == sorted(actual_params) -def test_update_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("reasoningEngine", ))) - - def test_update_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3585,17 +3403,20 @@ def test_delete_reasoning_engine_rest_required_fields(request_type=reasoning_eng # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteReasoningEngine, + "_BaseDeleteReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_reasoning_engine._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -3642,13 +3463,6 @@ def test_delete_reasoning_engine_rest_required_fields(request_type=reasoning_eng assert sorted(expected_params) == sorted(actual_params) -def test_delete_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_schedule_service.py b/tests/unit/gapic/aiplatform_v1/test_schedule_service.py index 832b60aa91..9772c907d6 100644 --- a/tests/unit/gapic/aiplatform_v1/test_schedule_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_schedule_service.py @@ -110,6 +110,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -153,150 +165,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ScheduleServiceClient._get_default_mtls_endpoint(None) is None - assert ScheduleServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ScheduleServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ScheduleServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ScheduleServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ScheduleServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ScheduleServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ScheduleServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ScheduleServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -310,40 +178,6 @@ def test__get_client_cert_source(): assert ScheduleServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ScheduleServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ScheduleServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ScheduleServiceClient)) -@mock.patch.object(ScheduleServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ScheduleServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ScheduleServiceClient._DEFAULT_UNIVERSE - default_endpoint = ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ScheduleServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ScheduleServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - assert ScheduleServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ScheduleServiceClient._get_api_endpoint(None, None, default_universe, "always") == ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - assert ScheduleServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - assert ScheduleServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ScheduleServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ScheduleServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ScheduleServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ScheduleServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ScheduleServiceClient._get_universe_domain(None, None) == ScheduleServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ScheduleServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -728,11 +562,12 @@ def test_schedule_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -775,11 +610,12 @@ def test_schedule_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3551,16 +3387,18 @@ def test_create_schedule_rest_required_fields(request_type=schedule_service.Crea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateSchedule, + "_BaseCreateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3610,13 +3448,6 @@ def test_create_schedule_rest_required_fields(request_type=schedule_service.Crea assert sorted(expected_params) == sorted(actual_params) -def test_create_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "schedule", ))) - - def test_create_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3725,16 +3556,18 @@ def test_delete_schedule_rest_required_fields(request_type=schedule_service.Dele # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSchedule, + "_BaseDeleteSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3780,13 +3613,6 @@ def test_delete_schedule_rest_required_fields(request_type=schedule_service.Dele assert sorted(expected_params) == sorted(actual_params) -def test_delete_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3887,16 +3713,18 @@ def test_get_schedule_rest_required_fields(request_type=schedule_service.GetSche # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetSchedule, + "_BaseGetSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3945,13 +3773,6 @@ def test_get_schedule_rest_required_fields(request_type=schedule_service.GetSche assert sorted(expected_params) == sorted(actual_params) -def test_get_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4054,17 +3875,20 @@ def test_list_schedules_rest_required_fields(request_type=schedule_service.ListS # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_schedules._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSchedules, + "_BaseListSchedules__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_schedules._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4114,13 +3938,6 @@ def test_list_schedules_rest_required_fields(request_type=schedule_service.ListS assert sorted(expected_params) == sorted(actual_params) -def test_list_schedules_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_schedules._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_schedules_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4288,16 +4105,18 @@ def test_pause_schedule_rest_required_fields(request_type=schedule_service.Pause # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePauseSchedule, + "_BasePauseSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4344,13 +4163,6 @@ def test_pause_schedule_rest_required_fields(request_type=schedule_service.Pause assert sorted(expected_params) == sorted(actual_params) -def test_pause_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.pause_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_pause_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4451,16 +4263,18 @@ def test_resume_schedule_rest_required_fields(request_type=schedule_service.Resu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseResumeSchedule, + "_BaseResumeSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4507,13 +4321,6 @@ def test_resume_schedule_rest_required_fields(request_type=schedule_service.Resu assert sorted(expected_params) == sorted(actual_params) -def test_resume_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.resume_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_resume_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4615,15 +4422,18 @@ def test_update_schedule_rest_required_fields(request_type=schedule_service.Upda # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateSchedule, + "_BaseUpdateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_schedule._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4672,13 +4482,6 @@ def test_update_schedule_rest_required_fields(request_type=schedule_service.Upda assert sorted(expected_params) == sorted(actual_params) -def test_update_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("schedule", "updateMask", ))) - - def test_update_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_session_service.py b/tests/unit/gapic/aiplatform_v1/test_session_service.py index a01acf7c93..12f64bebd0 100644 --- a/tests/unit/gapic/aiplatform_v1/test_session_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_session_service.py @@ -94,6 +94,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -137,150 +149,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert SessionServiceClient._get_default_mtls_endpoint(None) is None - assert SessionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert SessionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert SessionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SessionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SessionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - SessionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert SessionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert SessionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert SessionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert SessionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - SessionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert SessionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - SessionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert SessionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -294,40 +162,6 @@ def test__get_client_cert_source(): assert SessionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert SessionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(SessionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SessionServiceClient)) -@mock.patch.object(SessionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SessionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = SessionServiceClient._DEFAULT_UNIVERSE - default_endpoint = SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert SessionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert SessionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == SessionServiceClient.DEFAULT_MTLS_ENDPOINT - assert SessionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert SessionServiceClient._get_api_endpoint(None, None, default_universe, "always") == SessionServiceClient.DEFAULT_MTLS_ENDPOINT - assert SessionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == SessionServiceClient.DEFAULT_MTLS_ENDPOINT - assert SessionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert SessionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - SessionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert SessionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert SessionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert SessionServiceClient._get_universe_domain(None, None) == SessionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - SessionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -712,11 +546,12 @@ def test_session_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -759,11 +594,12 @@ def test_session_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3689,17 +3525,20 @@ def test_create_session_rest_required_fields(request_type=session_service.Create # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateSession, + "_BaseCreateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_session._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("session_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("sessionId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3747,13 +3586,6 @@ def test_create_session_rest_required_fields(request_type=session_service.Create assert sorted(expected_params) == sorted(actual_params) -def test_create_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(("sessionId", )) & set(("parent", "session", ))) - - def test_create_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3858,16 +3690,18 @@ def test_get_session_rest_required_fields(request_type=session_service.GetSessio # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetSession, + "_BaseGetSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_session._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3916,13 +3750,6 @@ def test_get_session_rest_required_fields(request_type=session_service.GetSessio assert sorted(expected_params) == sorted(actual_params) -def test_get_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4025,17 +3852,20 @@ def test_list_sessions_rest_required_fields(request_type=session_service.ListSes # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_sessions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSessions, + "_BaseListSessions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_sessions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4085,13 +3915,6 @@ def test_list_sessions_rest_required_fields(request_type=session_service.ListSes assert sorted(expected_params) == sorted(actual_params) -def test_list_sessions_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_sessions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_sessions_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4258,15 +4081,18 @@ def test_update_session_rest_required_fields(request_type=session_service.Update # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateSession, + "_BaseUpdateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_session._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4315,13 +4141,6 @@ def test_update_session_rest_required_fields(request_type=session_service.Update assert sorted(expected_params) == sorted(actual_params) -def test_update_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("session", ))) - - def test_update_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4430,16 +4249,18 @@ def test_delete_session_rest_required_fields(request_type=session_service.Delete # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSession, + "_BaseDeleteSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_session._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4485,13 +4306,6 @@ def test_delete_session_rest_required_fields(request_type=session_service.Delete assert sorted(expected_params) == sorted(actual_params) -def test_delete_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4592,17 +4406,20 @@ def test_list_events_rest_required_fields(request_type=session_service.ListEvent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_events._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListEvents, + "_BaseListEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_events._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4652,13 +4469,6 @@ def test_list_events_rest_required_fields(request_type=session_service.ListEvent assert sorted(expected_params) == sorted(actual_params) -def test_list_events_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_events._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_events_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4826,16 +4636,18 @@ def test_append_event_rest_required_fields(request_type=session_service.AppendEv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).append_event._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAppendEvent, + "_BaseAppendEvent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).append_event._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4885,13 +4697,6 @@ def test_append_event_rest_required_fields(request_type=session_service.AppendEv assert sorted(expected_params) == sorted(actual_params) -def test_append_event_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.append_event._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "event", ))) - - def test_append_event_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_specialist_pool_service.py b/tests/unit/gapic/aiplatform_v1/test_specialist_pool_service.py index 87b7ba17a4..32b2ef1f89 100644 --- a/tests/unit/gapic/aiplatform_v1/test_specialist_pool_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_specialist_pool_service.py @@ -88,6 +88,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -131,150 +143,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(None) is None - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - SpecialistPoolServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert SpecialistPoolServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - SpecialistPoolServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - SpecialistPoolServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -288,40 +156,6 @@ def test__get_client_cert_source(): assert SpecialistPoolServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert SpecialistPoolServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(SpecialistPoolServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SpecialistPoolServiceClient)) -@mock.patch.object(SpecialistPoolServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SpecialistPoolServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = SpecialistPoolServiceClient._DEFAULT_UNIVERSE - default_endpoint = SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert SpecialistPoolServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert SpecialistPoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, default_universe, "always") == SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert SpecialistPoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - SpecialistPoolServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert SpecialistPoolServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert SpecialistPoolServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert SpecialistPoolServiceClient._get_universe_domain(None, None) == SpecialistPoolServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - SpecialistPoolServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -706,11 +540,12 @@ def test_specialist_pool_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -753,11 +588,12 @@ def test_specialist_pool_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2850,16 +2686,18 @@ def test_create_specialist_pool_rest_required_fields(request_type=specialist_poo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateSpecialistPool, + "_BaseCreateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_specialist_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2906,13 +2744,6 @@ def test_create_specialist_pool_rest_required_fields(request_type=specialist_poo assert sorted(expected_params) == sorted(actual_params) -def test_create_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "specialistPool", ))) - - def test_create_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3015,16 +2846,18 @@ def test_get_specialist_pool_rest_required_fields(request_type=specialist_pool_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetSpecialistPool, + "_BaseGetSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_specialist_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3073,13 +2906,6 @@ def test_get_specialist_pool_rest_required_fields(request_type=specialist_pool_s assert sorted(expected_params) == sorted(actual_params) -def test_get_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3182,17 +3008,20 @@ def test_list_specialist_pools_rest_required_fields(request_type=specialist_pool # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_specialist_pools._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSpecialistPools, + "_BaseListSpecialistPools__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_specialist_pools._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3242,13 +3071,6 @@ def test_list_specialist_pools_rest_required_fields(request_type=specialist_pool assert sorted(expected_params) == sorted(actual_params) -def test_list_specialist_pools_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_specialist_pools._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_specialist_pools_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3420,17 +3242,20 @@ def test_delete_specialist_pool_rest_required_fields(request_type=specialist_poo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSpecialistPool, + "_BaseDeleteSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_specialist_pool._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -3477,13 +3302,6 @@ def test_delete_specialist_pool_rest_required_fields(request_type=specialist_poo assert sorted(expected_params) == sorted(actual_params) -def test_delete_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3587,15 +3405,18 @@ def test_update_specialist_pool_rest_required_fields(request_type=specialist_poo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateSpecialistPool, + "_BaseUpdateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_specialist_pool._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3641,13 +3462,6 @@ def test_update_specialist_pool_rest_required_fields(request_type=specialist_poo assert sorted(expected_params) == sorted(actual_params) -def test_update_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("specialistPool", "updateMask", ))) - - def test_update_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_tensorboard_service.py b/tests/unit/gapic/aiplatform_v1/test_tensorboard_service.py index 44b3d011a9..e2e4e453e9 100644 --- a/tests/unit/gapic/aiplatform_v1/test_tensorboard_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_tensorboard_service.py @@ -97,6 +97,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -140,150 +152,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert TensorboardServiceClient._get_default_mtls_endpoint(None) is None - assert TensorboardServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert TensorboardServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert TensorboardServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - TensorboardServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert TensorboardServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - TensorboardServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - TensorboardServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert TensorboardServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -297,40 +165,6 @@ def test__get_client_cert_source(): assert TensorboardServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert TensorboardServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(TensorboardServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(TensorboardServiceClient)) -@mock.patch.object(TensorboardServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(TensorboardServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = TensorboardServiceClient._DEFAULT_UNIVERSE - default_endpoint = TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert TensorboardServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert TensorboardServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - assert TensorboardServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert TensorboardServiceClient._get_api_endpoint(None, None, default_universe, "always") == TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - assert TensorboardServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - assert TensorboardServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert TensorboardServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - TensorboardServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert TensorboardServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert TensorboardServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert TensorboardServiceClient._get_universe_domain(None, None) == TensorboardServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - TensorboardServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -715,11 +549,12 @@ def test_tensorboard_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -762,11 +597,12 @@ def test_tensorboard_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -11694,16 +11530,18 @@ def test_create_tensorboard_rest_required_fields(request_type=tensorboard_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboard, + "_BaseCreateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -11750,13 +11588,6 @@ def test_create_tensorboard_rest_required_fields(request_type=tensorboard_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "tensorboard", ))) - - def test_create_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11859,16 +11690,18 @@ def test_get_tensorboard_rest_required_fields(request_type=tensorboard_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboard, + "_BaseGetTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11917,13 +11750,6 @@ def test_get_tensorboard_rest_required_fields(request_type=tensorboard_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12029,15 +11855,18 @@ def test_update_tensorboard_rest_required_fields(request_type=tensorboard_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboard, + "_BaseUpdateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -12083,13 +11912,6 @@ def test_update_tensorboard_rest_required_fields(request_type=tensorboard_servic assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboard", ))) - - def test_update_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12192,17 +12014,20 @@ def test_list_tensorboards_rest_required_fields(request_type=tensorboard_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboards._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboards, + "_BaseListTensorboards__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboards._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12252,13 +12077,6 @@ def test_list_tensorboards_rest_required_fields(request_type=tensorboard_service assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboards_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboards._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboards_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12430,16 +12248,18 @@ def test_delete_tensorboard_rest_required_fields(request_type=tensorboard_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboard, + "_BaseDeleteTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -12485,13 +12305,6 @@ def test_delete_tensorboard_rest_required_fields(request_type=tensorboard_servic assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12592,16 +12405,18 @@ def test_read_tensorboard_usage_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_usage._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardUsage, + "_BaseReadTensorboardUsage__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboard"] = 'tensorboard_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_usage._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboard" in jsonified_request assert jsonified_request["tensorboard"] == 'tensorboard_value' @@ -12650,13 +12465,6 @@ def test_read_tensorboard_usage_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_usage_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_usage._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboard", ))) - - def test_read_tensorboard_usage_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12759,16 +12567,18 @@ def test_read_tensorboard_size_rest_required_fields(request_type=tensorboard_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_size._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardSize, + "_BaseReadTensorboardSize__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboard"] = 'tensorboard_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_size._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboard" in jsonified_request assert jsonified_request["tensorboard"] == 'tensorboard_value' @@ -12817,13 +12627,6 @@ def test_read_tensorboard_size_rest_required_fields(request_type=tensorboard_ser assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_size_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_size._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboard", ))) - - def test_read_tensorboard_size_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12928,7 +12731,12 @@ def test_create_tensorboard_experiment_rest_required_fields(request_type=tensorb # verify fields with default values are dropped assert "tensorboardExperimentId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboardExperiment, + "_BaseCreateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -12938,10 +12746,8 @@ def test_create_tensorboard_experiment_rest_required_fields(request_type=tensorb jsonified_request["parent"] = 'parent_value' jsonified_request["tensorboardExperimentId"] = 'tensorboard_experiment_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_experiment._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("tensorboard_experiment_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("tensorboardExperimentId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12998,13 +12804,6 @@ def test_create_tensorboard_experiment_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(("tensorboardExperimentId", )) & set(("parent", "tensorboardExperimentId", ))) - - def test_create_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13111,16 +12910,18 @@ def test_get_tensorboard_experiment_rest_required_fields(request_type=tensorboar # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboardExperiment, + "_BaseGetTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_experiment._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13169,13 +12970,6 @@ def test_get_tensorboard_experiment_rest_required_fields(request_type=tensorboar assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13277,15 +13071,18 @@ def test_update_tensorboard_experiment_rest_required_fields(request_type=tensorb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboardExperiment, + "_BaseUpdateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_experiment._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -13334,13 +13131,6 @@ def test_update_tensorboard_experiment_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboardExperiment", ))) - - def test_update_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13445,17 +13235,20 @@ def test_list_tensorboard_experiments_rest_required_fields(request_type=tensorbo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_experiments._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboardExperiments, + "_BaseListTensorboardExperiments__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_experiments._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13505,13 +13298,6 @@ def test_list_tensorboard_experiments_rest_required_fields(request_type=tensorbo assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboard_experiments_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboard_experiments._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboard_experiments_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13683,16 +13469,18 @@ def test_delete_tensorboard_experiment_rest_required_fields(request_type=tensorb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboardExperiment, + "_BaseDeleteTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_experiment._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13738,13 +13526,6 @@ def test_delete_tensorboard_experiment_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13847,7 +13628,12 @@ def test_create_tensorboard_run_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped assert "tensorboardRunId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboardRun, + "_BaseCreateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -13857,10 +13643,8 @@ def test_create_tensorboard_run_rest_required_fields(request_type=tensorboard_se jsonified_request["parent"] = 'parent_value' jsonified_request["tensorboardRunId"] = 'tensorboard_run_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_run._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("tensorboard_run_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("tensorboardRunId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13917,13 +13701,6 @@ def test_create_tensorboard_run_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(("tensorboardRunId", )) & set(("parent", "tensorboardRun", "tensorboardRunId", ))) - - def test_create_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14030,16 +13807,18 @@ def test_batch_create_tensorboard_runs_rest_required_fields(request_type=tensorb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_runs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateTensorboardRuns, + "_BaseBatchCreateTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_runs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14089,13 +13868,6 @@ def test_batch_create_tensorboard_runs_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_tensorboard_runs_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_tensorboard_runs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_tensorboard_runs_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14200,16 +13972,18 @@ def test_get_tensorboard_run_rest_required_fields(request_type=tensorboard_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboardRun, + "_BaseGetTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_run._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14258,13 +14032,6 @@ def test_get_tensorboard_run_rest_required_fields(request_type=tensorboard_servi assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14366,15 +14133,18 @@ def test_update_tensorboard_run_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboardRun, + "_BaseUpdateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_run._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -14423,13 +14193,6 @@ def test_update_tensorboard_run_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboardRun", ))) - - def test_update_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14534,17 +14297,20 @@ def test_list_tensorboard_runs_rest_required_fields(request_type=tensorboard_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_runs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboardRuns, + "_BaseListTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_runs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14594,13 +14360,6 @@ def test_list_tensorboard_runs_rest_required_fields(request_type=tensorboard_ser assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboard_runs_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboard_runs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboard_runs_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14772,16 +14531,18 @@ def test_delete_tensorboard_run_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboardRun, + "_BaseDeleteTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_run._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14827,13 +14588,6 @@ def test_delete_tensorboard_run_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14934,16 +14688,18 @@ def test_batch_create_tensorboard_time_series_rest_required_fields(request_type= # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateTensorboardTimeSeries, + "_BaseBatchCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_time_series._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14993,13 +14749,6 @@ def test_batch_create_tensorboard_time_series_rest_required_fields(request_type= assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15104,17 +14853,20 @@ def test_create_tensorboard_time_series_rest_required_fields(request_type=tensor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboardTimeSeries, + "_BaseCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_time_series._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("tensorboard_time_series_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("tensorboardTimeSeriesId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -15165,13 +14917,6 @@ def test_create_tensorboard_time_series_rest_required_fields(request_type=tensor assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(("tensorboardTimeSeriesId", )) & set(("parent", "tensorboardTimeSeries", ))) - - def test_create_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15276,16 +15021,18 @@ def test_get_tensorboard_time_series_rest_required_fields(request_type=tensorboa # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboardTimeSeries, + "_BaseGetTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_time_series._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15334,13 +15081,6 @@ def test_get_tensorboard_time_series_rest_required_fields(request_type=tensorboa assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15442,15 +15182,18 @@ def test_update_tensorboard_time_series_rest_required_fields(request_type=tensor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboardTimeSeries, + "_BaseUpdateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_time_series._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -15499,13 +15242,6 @@ def test_update_tensorboard_time_series_rest_required_fields(request_type=tensor assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboardTimeSeries", ))) - - def test_update_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15610,17 +15346,20 @@ def test_list_tensorboard_time_series_rest_required_fields(request_type=tensorbo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboardTimeSeries, + "_BaseListTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_time_series._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -15670,13 +15409,6 @@ def test_list_tensorboard_time_series_rest_required_fields(request_type=tensorbo assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15848,16 +15580,18 @@ def test_delete_tensorboard_time_series_rest_required_fields(request_type=tensor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboardTimeSeries, + "_BaseDeleteTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_time_series._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15903,13 +15637,6 @@ def test_delete_tensorboard_time_series_rest_required_fields(request_type=tensor assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16012,7 +15739,12 @@ def test_batch_read_tensorboard_time_series_data_rest_required_fields(request_ty # verify fields with default values are dropped assert "timeSeries" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchReadTensorboardTimeSeriesData, + "_BaseBatchReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -16022,10 +15754,8 @@ def test_batch_read_tensorboard_time_series_data_rest_required_fields(request_ty jsonified_request["tensorboard"] = 'tensorboard_value' jsonified_request["timeSeries"] = 'time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("time_series", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("timeSeries", )) # verify required fields with non-default values are left alone assert "tensorboard" in jsonified_request @@ -16081,13 +15811,6 @@ def test_batch_read_tensorboard_time_series_data_rest_required_fields(request_ty assert sorted(expected_params) == sorted(actual_params) -def test_batch_read_tensorboard_time_series_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_read_tensorboard_time_series_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(("timeSeries", )) & set(("tensorboard", "timeSeries", ))) - - def test_batch_read_tensorboard_time_series_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16190,17 +15913,20 @@ def test_read_tensorboard_time_series_data_rest_required_fields(request_type=ten # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardTimeSeriesData, + "_BaseReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardTimeSeries"] = 'tensorboard_time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "max_data_points", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "maxDataPoints", )) # verify required fields with non-default values are left alone assert "tensorboardTimeSeries" in jsonified_request @@ -16250,13 +15976,6 @@ def test_read_tensorboard_time_series_data_rest_required_fields(request_type=ten assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_time_series_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_time_series_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "maxDataPoints", )) & set(("tensorboardTimeSeries", ))) - - def test_read_tensorboard_time_series_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16359,17 +16078,20 @@ def test_read_tensorboard_blob_data_rest_required_fields(request_type=tensorboar # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_blob_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardBlobData, + "_BaseReadTensorboardBlobData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["timeSeries"] = 'time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_blob_data._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("blob_ids", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("blobIds", )) # verify required fields with non-default values are left alone assert "timeSeries" in jsonified_request @@ -16422,13 +16144,6 @@ def test_read_tensorboard_blob_data_rest_required_fields(request_type=tensorboar assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_blob_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_blob_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(("blobIds", )) & set(("timeSeries", ))) - - def test_read_tensorboard_blob_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16534,16 +16249,18 @@ def test_write_tensorboard_experiment_data_rest_required_fields(request_type=ten # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_experiment_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseWriteTensorboardExperimentData, + "_BaseWriteTensorboardExperimentData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardExperiment"] = 'tensorboard_experiment_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_experiment_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboardExperiment" in jsonified_request assert jsonified_request["tensorboardExperiment"] == 'tensorboard_experiment_value' @@ -16593,13 +16310,6 @@ def test_write_tensorboard_experiment_data_rest_required_fields(request_type=ten assert sorted(expected_params) == sorted(actual_params) -def test_write_tensorboard_experiment_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.write_tensorboard_experiment_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboardExperiment", "writeRunDataRequests", ))) - - def test_write_tensorboard_experiment_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16704,16 +16414,18 @@ def test_write_tensorboard_run_data_rest_required_fields(request_type=tensorboar # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_run_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseWriteTensorboardRunData, + "_BaseWriteTensorboardRunData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardRun"] = 'tensorboard_run_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_run_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboardRun" in jsonified_request assert jsonified_request["tensorboardRun"] == 'tensorboard_run_value' @@ -16763,13 +16475,6 @@ def test_write_tensorboard_run_data_rest_required_fields(request_type=tensorboar assert sorted(expected_params) == sorted(actual_params) -def test_write_tensorboard_run_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.write_tensorboard_run_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboardRun", "timeSeriesData", ))) - - def test_write_tensorboard_run_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16874,16 +16579,18 @@ def test_export_tensorboard_time_series_data_rest_required_fields(request_type=t # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportTensorboardTimeSeriesData, + "_BaseExportTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardTimeSeries"] = 'tensorboard_time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboardTimeSeries" in jsonified_request assert jsonified_request["tensorboardTimeSeries"] == 'tensorboard_time_series_value' @@ -16933,13 +16640,6 @@ def test_export_tensorboard_time_series_data_rest_required_fields(request_type=t assert sorted(expected_params) == sorted(actual_params) -def test_export_tensorboard_time_series_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_tensorboard_time_series_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboardTimeSeries", ))) - - def test_export_tensorboard_time_series_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_vertex_rag_data_service.py b/tests/unit/gapic/aiplatform_v1/test_vertex_rag_data_service.py index b37206eee4..172ecd6e4d 100644 --- a/tests/unit/gapic/aiplatform_v1/test_vertex_rag_data_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_vertex_rag_data_service.py @@ -91,6 +91,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -134,150 +146,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert VertexRagDataServiceClient._get_default_mtls_endpoint(None) is None - assert VertexRagDataServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert VertexRagDataServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagDataServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - VertexRagDataServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert VertexRagDataServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagDataServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - VertexRagDataServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -291,40 +159,6 @@ def test__get_client_cert_source(): assert VertexRagDataServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert VertexRagDataServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(VertexRagDataServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagDataServiceClient)) -@mock.patch.object(VertexRagDataServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagDataServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = VertexRagDataServiceClient._DEFAULT_UNIVERSE - default_endpoint = VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert VertexRagDataServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert VertexRagDataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagDataServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert VertexRagDataServiceClient._get_api_endpoint(None, None, default_universe, "always") == VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagDataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagDataServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert VertexRagDataServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagDataServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert VertexRagDataServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert VertexRagDataServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert VertexRagDataServiceClient._get_universe_domain(None, None) == VertexRagDataServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - VertexRagDataServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -709,11 +543,12 @@ def test_vertex_rag_data_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -756,11 +591,12 @@ def test_vertex_rag_data_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5251,16 +5087,18 @@ def test_create_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateRagCorpus, + "_BaseCreateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_corpus._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -5307,13 +5145,6 @@ def test_create_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser assert sorted(expected_params) == sorted(actual_params) -def test_create_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "ragCorpus", ))) - - def test_create_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5419,14 +5250,16 @@ def test_update_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateRagCorpus, + "_BaseUpdateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_corpus._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone client = VertexRagDataServiceClient( @@ -5471,13 +5304,6 @@ def test_update_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser assert sorted(expected_params) == sorted(actual_params) -def test_update_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("ragCorpus", ))) - - def test_update_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5578,16 +5404,18 @@ def test_get_rag_corpus_rest_required_fields(request_type=vertex_rag_data_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagCorpus, + "_BaseGetRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_corpus._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5636,13 +5464,6 @@ def test_get_rag_corpus_rest_required_fields(request_type=vertex_rag_data_servic assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5745,17 +5566,20 @@ def test_list_rag_corpora_rest_required_fields(request_type=vertex_rag_data_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_corpora._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListRagCorpora, + "_BaseListRagCorpora__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_corpora._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5805,13 +5629,6 @@ def test_list_rag_corpora_rest_required_fields(request_type=vertex_rag_data_serv assert sorted(expected_params) == sorted(actual_params) -def test_list_rag_corpora_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_rag_corpora._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_rag_corpora_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5983,17 +5800,20 @@ def test_delete_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteRagCorpus, + "_BaseDeleteRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_corpus._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -6040,13 +5860,6 @@ def test_delete_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser assert sorted(expected_params) == sorted(actual_params) -def test_delete_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6147,16 +5960,18 @@ def test_upload_rag_file_rest_required_fields(request_type=vertex_rag_data_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_rag_file._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUploadRagFile, + "_BaseUploadRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_rag_file._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6206,13 +6021,6 @@ def test_upload_rag_file_rest_required_fields(request_type=vertex_rag_data_servi assert sorted(expected_params) == sorted(actual_params) -def test_upload_rag_file_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upload_rag_file._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "ragFile", "uploadRagFileConfig", ))) - - def test_upload_rag_file_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6323,16 +6131,18 @@ def test_import_rag_files_rest_required_fields(request_type=vertex_rag_data_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_rag_files._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportRagFiles, + "_BaseImportRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_rag_files._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6379,13 +6189,6 @@ def test_import_rag_files_rest_required_fields(request_type=vertex_rag_data_serv assert sorted(expected_params) == sorted(actual_params) -def test_import_rag_files_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_rag_files._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "importRagFilesConfig", ))) - - def test_import_rag_files_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6488,16 +6291,18 @@ def test_get_rag_file_rest_required_fields(request_type=vertex_rag_data_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_file._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagFile, + "_BaseGetRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_file._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6546,13 +6351,6 @@ def test_get_rag_file_rest_required_fields(request_type=vertex_rag_data_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_file_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_file._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_file_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6655,17 +6453,20 @@ def test_list_rag_files_rest_required_fields(request_type=vertex_rag_data_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_files._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListRagFiles, + "_BaseListRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_files._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6715,13 +6516,6 @@ def test_list_rag_files_rest_required_fields(request_type=vertex_rag_data_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_rag_files_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_rag_files._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_rag_files_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6893,16 +6687,18 @@ def test_delete_rag_file_rest_required_fields(request_type=vertex_rag_data_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_file._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteRagFile, + "_BaseDeleteRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_file._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6948,13 +6744,6 @@ def test_delete_rag_file_rest_required_fields(request_type=vertex_rag_data_servi assert sorted(expected_params) == sorted(actual_params) -def test_delete_rag_file_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_rag_file._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_rag_file_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7058,14 +6847,16 @@ def test_update_rag_engine_config_rest_required_fields(request_type=vertex_rag_d # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_engine_config._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateRagEngineConfig, + "_BaseUpdateRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_engine_config._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone client = VertexRagDataServiceClient( @@ -7110,13 +6901,6 @@ def test_update_rag_engine_config_rest_required_fields(request_type=vertex_rag_d assert sorted(expected_params) == sorted(actual_params) -def test_update_rag_engine_config_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_rag_engine_config._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("ragEngineConfig", ))) - - def test_update_rag_engine_config_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7217,16 +7001,18 @@ def test_get_rag_engine_config_rest_required_fields(request_type=vertex_rag_data # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_engine_config._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagEngineConfig, + "_BaseGetRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_engine_config._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7275,13 +7061,6 @@ def test_get_rag_engine_config_rest_required_fields(request_type=vertex_rag_data assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_engine_config_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_engine_config._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_engine_config_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_vertex_rag_service.py b/tests/unit/gapic/aiplatform_v1/test_vertex_rag_service.py index 63d29d0d3f..8954261bb5 100644 --- a/tests/unit/gapic/aiplatform_v1/test_vertex_rag_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_vertex_rag_service.py @@ -89,6 +89,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -132,150 +144,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert VertexRagServiceClient._get_default_mtls_endpoint(None) is None - assert VertexRagServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert VertexRagServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - VertexRagServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert VertexRagServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - VertexRagServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert VertexRagServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -289,40 +157,6 @@ def test__get_client_cert_source(): assert VertexRagServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert VertexRagServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(VertexRagServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagServiceClient)) -@mock.patch.object(VertexRagServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = VertexRagServiceClient._DEFAULT_UNIVERSE - default_endpoint = VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert VertexRagServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert VertexRagServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert VertexRagServiceClient._get_api_endpoint(None, None, default_universe, "always") == VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert VertexRagServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert VertexRagServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert VertexRagServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert VertexRagServiceClient._get_universe_domain(None, None) == VertexRagServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - VertexRagServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -707,11 +541,12 @@ def test_vertex_rag_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -754,11 +589,12 @@ def test_vertex_rag_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2650,16 +2486,18 @@ def test_retrieve_contexts_rest_required_fields(request_type=vertex_rag_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).retrieve_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRetrieveContexts, + "_BaseRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).retrieve_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2709,13 +2547,6 @@ def test_retrieve_contexts_rest_required_fields(request_type=vertex_rag_service. assert sorted(expected_params) == sorted(actual_params) -def test_retrieve_contexts_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.retrieve_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "query", ))) - - def test_retrieve_contexts_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2820,16 +2651,18 @@ def test_augment_prompt_rest_required_fields(request_type=vertex_rag_service.Aug # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).augment_prompt._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAugmentPrompt, + "_BaseAugmentPrompt__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).augment_prompt._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2879,13 +2712,6 @@ def test_augment_prompt_rest_required_fields(request_type=vertex_rag_service.Aug assert sorted(expected_params) == sorted(actual_params) -def test_augment_prompt_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.augment_prompt._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_augment_prompt_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2991,16 +2817,18 @@ def test_corroborate_content_rest_required_fields(request_type=vertex_rag_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).corroborate_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCorroborateContent, + "_BaseCorroborateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).corroborate_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3050,13 +2878,6 @@ def test_corroborate_content_rest_required_fields(request_type=vertex_rag_servic assert sorted(expected_params) == sorted(actual_params) -def test_corroborate_content_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.corroborate_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_corroborate_content_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3163,16 +2984,18 @@ def test_ask_contexts_rest_required_fields(request_type=vertex_rag_service.AskCo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).ask_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAskContexts, + "_BaseAskContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).ask_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3222,13 +3045,6 @@ def test_ask_contexts_rest_required_fields(request_type=vertex_rag_service.AskCo assert sorted(expected_params) == sorted(actual_params) -def test_ask_contexts_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.ask_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "query", ))) - - def test_ask_contexts_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3337,16 +3153,18 @@ def test_async_retrieve_contexts_rest_required_fields(request_type=vertex_rag_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_retrieve_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAsyncRetrieveContexts, + "_BaseAsyncRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_retrieve_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3393,13 +3211,6 @@ def test_async_retrieve_contexts_rest_required_fields(request_type=vertex_rag_se assert sorted(expected_params) == sorted(actual_params) -def test_async_retrieve_contexts_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.async_retrieve_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "query", ))) - - def test_async_retrieve_contexts_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1/test_vizier_service.py b/tests/unit/gapic/aiplatform_v1/test_vizier_service.py index fbc9e874c6..983739bd9e 100644 --- a/tests/unit/gapic/aiplatform_v1/test_vizier_service.py +++ b/tests/unit/gapic/aiplatform_v1/test_vizier_service.py @@ -89,6 +89,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -132,150 +144,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert VizierServiceClient._get_default_mtls_endpoint(None) is None - assert VizierServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert VizierServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert VizierServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VizierServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VizierServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - VizierServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert VizierServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert VizierServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert VizierServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert VizierServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - VizierServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert VizierServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - VizierServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert VizierServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -289,40 +157,6 @@ def test__get_client_cert_source(): assert VizierServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert VizierServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(VizierServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VizierServiceClient)) -@mock.patch.object(VizierServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VizierServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = VizierServiceClient._DEFAULT_UNIVERSE - default_endpoint = VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert VizierServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert VizierServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == VizierServiceClient.DEFAULT_MTLS_ENDPOINT - assert VizierServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert VizierServiceClient._get_api_endpoint(None, None, default_universe, "always") == VizierServiceClient.DEFAULT_MTLS_ENDPOINT - assert VizierServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == VizierServiceClient.DEFAULT_MTLS_ENDPOINT - assert VizierServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert VizierServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - VizierServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert VizierServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert VizierServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert VizierServiceClient._get_universe_domain(None, None) == VizierServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - VizierServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -707,11 +541,12 @@ def test_vizier_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -754,11 +589,12 @@ def test_vizier_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5824,16 +5660,18 @@ def test_create_study_rest_required_fields(request_type=vizier_service.CreateStu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateStudy, + "_BaseCreateStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -5883,13 +5721,6 @@ def test_create_study_rest_required_fields(request_type=vizier_service.CreateStu assert sorted(expected_params) == sorted(actual_params) -def test_create_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "study", ))) - - def test_create_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5994,16 +5825,18 @@ def test_get_study_rest_required_fields(request_type=vizier_service.GetStudyRequ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetStudy, + "_BaseGetStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6052,13 +5885,6 @@ def test_get_study_rest_required_fields(request_type=vizier_service.GetStudyRequ assert sorted(expected_params) == sorted(actual_params) -def test_get_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6161,17 +5987,20 @@ def test_list_studies_rest_required_fields(request_type=vizier_service.ListStudi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_studies._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListStudies, + "_BaseListStudies__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_studies._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6221,13 +6050,6 @@ def test_list_studies_rest_required_fields(request_type=vizier_service.ListStudi assert sorted(expected_params) == sorted(actual_params) -def test_list_studies_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_studies._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_studies_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6395,16 +6217,18 @@ def test_delete_study_rest_required_fields(request_type=vizier_service.DeleteStu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteStudy, + "_BaseDeleteStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6450,13 +6274,6 @@ def test_delete_study_rest_required_fields(request_type=vizier_service.DeleteStu assert sorted(expected_params) == sorted(actual_params) -def test_delete_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6558,7 +6375,12 @@ def test_lookup_study_rest_required_fields(request_type=vizier_service.LookupStu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).lookup_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseLookupStudy, + "_BaseLookupStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -6566,9 +6388,6 @@ def test_lookup_study_rest_required_fields(request_type=vizier_service.LookupStu jsonified_request["parent"] = 'parent_value' jsonified_request["displayName"] = 'display_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).lookup_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6620,13 +6439,6 @@ def test_lookup_study_rest_required_fields(request_type=vizier_service.LookupStu assert sorted(expected_params) == sorted(actual_params) -def test_lookup_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.lookup_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "displayName", ))) - - def test_lookup_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6735,7 +6547,12 @@ def test_suggest_trials_rest_required_fields(request_type=vizier_service.Suggest # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).suggest_trials._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSuggestTrials, + "_BaseSuggestTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -6744,9 +6561,6 @@ def test_suggest_trials_rest_required_fields(request_type=vizier_service.Suggest jsonified_request["suggestionCount"] = 1744 jsonified_request["clientId"] = 'client_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).suggest_trials._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6797,13 +6611,6 @@ def test_suggest_trials_rest_required_fields(request_type=vizier_service.Suggest assert sorted(expected_params) == sorted(actual_params) -def test_suggest_trials_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.suggest_trials._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "suggestionCount", "clientId", ))) - - def test_create_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -6852,16 +6659,18 @@ def test_create_trial_rest_required_fields(request_type=vizier_service.CreateTri # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTrial, + "_BaseCreateTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6911,13 +6720,6 @@ def test_create_trial_rest_required_fields(request_type=vizier_service.CreateTri assert sorted(expected_params) == sorted(actual_params) -def test_create_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "trial", ))) - - def test_create_trial_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7022,16 +6824,18 @@ def test_get_trial_rest_required_fields(request_type=vizier_service.GetTrialRequ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTrial, + "_BaseGetTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7080,13 +6884,6 @@ def test_get_trial_rest_required_fields(request_type=vizier_service.GetTrialRequ assert sorted(expected_params) == sorted(actual_params) -def test_get_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_trial_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7189,17 +6986,20 @@ def test_list_trials_rest_required_fields(request_type=vizier_service.ListTrials # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_trials._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTrials, + "_BaseListTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_trials._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -7249,13 +7049,6 @@ def test_list_trials_rest_required_fields(request_type=vizier_service.ListTrials assert sorted(expected_params) == sorted(actual_params) -def test_list_trials_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_trials._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_trials_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7423,16 +7216,18 @@ def test_add_trial_measurement_rest_required_fields(request_type=vizier_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_trial_measurement._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddTrialMeasurement, + "_BaseAddTrialMeasurement__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["trialName"] = 'trial_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_trial_measurement._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "trialName" in jsonified_request assert jsonified_request["trialName"] == 'trial_name_value' @@ -7482,13 +7277,6 @@ def test_add_trial_measurement_rest_required_fields(request_type=vizier_service. assert sorted(expected_params) == sorted(actual_params) -def test_add_trial_measurement_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_trial_measurement._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("trialName", "measurement", ))) - - def test_complete_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -7537,16 +7325,18 @@ def test_complete_trial_rest_required_fields(request_type=vizier_service.Complet # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).complete_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCompleteTrial, + "_BaseCompleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).complete_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7596,13 +7386,6 @@ def test_complete_trial_rest_required_fields(request_type=vizier_service.Complet assert sorted(expected_params) == sorted(actual_params) -def test_complete_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.complete_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -7651,16 +7434,18 @@ def test_delete_trial_rest_required_fields(request_type=vizier_service.DeleteTri # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTrial, + "_BaseDeleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7706,13 +7491,6 @@ def test_delete_trial_rest_required_fields(request_type=vizier_service.DeleteTri assert sorted(expected_params) == sorted(actual_params) -def test_delete_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_trial_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7817,16 +7595,18 @@ def test_check_trial_early_stopping_state_rest_required_fields(request_type=vizi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).check_trial_early_stopping_state._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCheckTrialEarlyStoppingState, + "_BaseCheckTrialEarlyStoppingState__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["trialName"] = 'trial_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).check_trial_early_stopping_state._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "trialName" in jsonified_request assert jsonified_request["trialName"] == 'trial_name_value' @@ -7873,13 +7653,6 @@ def test_check_trial_early_stopping_state_rest_required_fields(request_type=vizi assert sorted(expected_params) == sorted(actual_params) -def test_check_trial_early_stopping_state_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.check_trial_early_stopping_state._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("trialName", ))) - - def test_stop_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -7928,16 +7701,18 @@ def test_stop_trial_rest_required_fields(request_type=vizier_service.StopTrialRe # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStopTrial, + "_BaseStopTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7987,13 +7762,6 @@ def test_stop_trial_rest_required_fields(request_type=vizier_service.StopTrialRe assert sorted(expected_params) == sorted(actual_params) -def test_stop_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stop_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_list_optimal_trials_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -8042,16 +7810,18 @@ def test_list_optimal_trials_rest_required_fields(request_type=vizier_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_optimal_trials._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListOptimalTrials, + "_BaseListOptimalTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_optimal_trials._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -8101,13 +7871,6 @@ def test_list_optimal_trials_rest_required_fields(request_type=vizier_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_optimal_trials_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_optimal_trials._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_list_optimal_trials_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_compat.py b/tests/unit/gapic/aiplatform_v1beta1/test_compat.py new file mode 100644 index 0000000000..fdf4222244 --- /dev/null +++ b/tests/unit/gapic/aiplatform_v1beta1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform_v1beta1._compat import transcode_request +from google.cloud.aiplatform_v1beta1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_dataset_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_dataset_service.py index 7418d9f112..181703bc61 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_dataset_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_dataset_service.py @@ -103,6 +103,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -146,150 +158,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert DatasetServiceClient._get_default_mtls_endpoint(None) is None - assert DatasetServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert DatasetServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert DatasetServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DatasetServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - DatasetServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert DatasetServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert DatasetServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert DatasetServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - DatasetServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert DatasetServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert DatasetServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - DatasetServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert DatasetServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert DatasetServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -303,40 +171,6 @@ def test__get_client_cert_source(): assert DatasetServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert DatasetServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(DatasetServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DatasetServiceClient)) -@mock.patch.object(DatasetServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DatasetServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = DatasetServiceClient._DEFAULT_UNIVERSE - default_endpoint = DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = DatasetServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert DatasetServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert DatasetServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - assert DatasetServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert DatasetServiceClient._get_api_endpoint(None, None, default_universe, "always") == DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - assert DatasetServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == DatasetServiceClient.DEFAULT_MTLS_ENDPOINT - assert DatasetServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert DatasetServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - DatasetServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert DatasetServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert DatasetServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert DatasetServiceClient._get_universe_domain(None, None) == DatasetServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - DatasetServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -721,11 +555,12 @@ def test_dataset_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -768,11 +603,12 @@ def test_dataset_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -8798,16 +8634,18 @@ def test_create_dataset_rest_required_fields(request_type=dataset_service.Create # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDataset, + "_BaseCreateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -8854,13 +8692,6 @@ def test_create_dataset_rest_required_fields(request_type=dataset_service.Create assert sorted(expected_params) == sorted(actual_params) -def test_create_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "dataset", ))) - - def test_create_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8963,17 +8794,20 @@ def test_get_dataset_rest_required_fields(request_type=dataset_service.GetDatase # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDataset, + "_BaseGetDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9023,13 +8857,6 @@ def test_get_dataset_rest_required_fields(request_type=dataset_service.GetDatase assert sorted(expected_params) == sorted(actual_params) -def test_get_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(("readMask", )) & set(("name", ))) - - def test_get_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9131,15 +8958,18 @@ def test_update_dataset_rest_required_fields(request_type=dataset_service.Update # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateDataset, + "_BaseUpdateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -9188,13 +9018,6 @@ def test_update_dataset_rest_required_fields(request_type=dataset_service.Update assert sorted(expected_params) == sorted(actual_params) -def test_update_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("dataset", "updateMask", ))) - - def test_update_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9299,17 +9122,20 @@ def test_list_datasets_rest_required_fields(request_type=dataset_service.ListDat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_datasets._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDatasets, + "_BaseListDatasets__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_datasets._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9359,13 +9185,6 @@ def test_list_datasets_rest_required_fields(request_type=dataset_service.ListDat assert sorted(expected_params) == sorted(actual_params) -def test_list_datasets_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_datasets._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_datasets_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9537,16 +9356,18 @@ def test_delete_dataset_rest_required_fields(request_type=dataset_service.Delete # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDataset, + "_BaseDeleteDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9592,13 +9413,6 @@ def test_delete_dataset_rest_required_fields(request_type=dataset_service.Delete assert sorted(expected_params) == sorted(actual_params) -def test_delete_dataset_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_dataset_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9703,16 +9517,18 @@ def test_import_data_rest_required_fields(request_type=dataset_service.ImportDat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportData, + "_BaseImportData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9759,13 +9575,6 @@ def test_import_data_rest_required_fields(request_type=dataset_service.ImportDat assert sorted(expected_params) == sorted(actual_params) -def test_import_data_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "importConfigs", ))) - - def test_import_data_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9872,16 +9681,18 @@ def test_export_data_rest_required_fields(request_type=dataset_service.ExportDat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportData, + "_BaseExportData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9928,13 +9739,6 @@ def test_export_data_rest_required_fields(request_type=dataset_service.ExportDat assert sorted(expected_params) == sorted(actual_params) -def test_export_data_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "exportConfig", ))) - - def test_export_data_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10041,16 +9845,18 @@ def test_create_dataset_version_rest_required_fields(request_type=dataset_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDatasetVersion, + "_BaseCreateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_dataset_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10097,13 +9903,6 @@ def test_create_dataset_version_rest_required_fields(request_type=dataset_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "datasetVersion", ))) - - def test_create_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10205,15 +10004,18 @@ def test_update_dataset_version_rest_required_fields(request_type=dataset_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateDatasetVersion, + "_BaseUpdateDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_dataset_version._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -10262,13 +10064,6 @@ def test_update_dataset_version_rest_required_fields(request_type=dataset_servic assert sorted(expected_params) == sorted(actual_params) -def test_update_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("datasetVersion", "updateMask", ))) - - def test_update_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10377,16 +10172,18 @@ def test_delete_dataset_version_rest_required_fields(request_type=dataset_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDatasetVersion, + "_BaseDeleteDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_dataset_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10432,13 +10229,6 @@ def test_delete_dataset_version_rest_required_fields(request_type=dataset_servic assert sorted(expected_params) == sorted(actual_params) -def test_delete_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10539,17 +10329,20 @@ def test_get_dataset_version_rest_required_fields(request_type=dataset_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDatasetVersion, + "_BaseGetDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_dataset_version._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -10599,13 +10392,6 @@ def test_get_dataset_version_rest_required_fields(request_type=dataset_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(("readMask", )) & set(("name", ))) - - def test_get_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10708,17 +10494,20 @@ def test_list_dataset_versions_rest_required_fields(request_type=dataset_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_dataset_versions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDatasetVersions, + "_BaseListDatasetVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_dataset_versions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10768,13 +10557,6 @@ def test_list_dataset_versions_rest_required_fields(request_type=dataset_service assert sorted(expected_params) == sorted(actual_params) -def test_list_dataset_versions_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_dataset_versions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_dataset_versions_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10946,16 +10728,18 @@ def test_restore_dataset_version_rest_required_fields(request_type=dataset_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).restore_dataset_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRestoreDatasetVersion, + "_BaseRestoreDatasetVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).restore_dataset_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11001,13 +10785,6 @@ def test_restore_dataset_version_rest_required_fields(request_type=dataset_servi assert sorted(expected_params) == sorted(actual_params) -def test_restore_dataset_version_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.restore_dataset_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_restore_dataset_version_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11108,17 +10885,20 @@ def test_list_data_items_rest_required_fields(request_type=dataset_service.ListD # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_items._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDataItems, + "_BaseListDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_items._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11168,13 +10948,6 @@ def test_list_data_items_rest_required_fields(request_type=dataset_service.ListD assert sorted(expected_params) == sorted(actual_params) -def test_list_data_items_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_data_items._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_data_items_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11342,17 +11115,20 @@ def test_search_data_items_rest_required_fields(request_type=dataset_service.Sea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_data_items._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchDataItems, + "_BaseSearchDataItems__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["dataset"] = 'dataset_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_data_items._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("annotation_filters", "annotations_filter", "annotations_limit", "data_item_filter", "data_labeling_job", "field_mask", "order_by", "order_by_annotation", "order_by_data_item", "page_size", "page_token", "saved_query", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("annotationFilters", "annotationsFilter", "annotationsLimit", "dataItemFilter", "dataLabelingJob", "fieldMask", "orderBy", "orderByAnnotation", "orderByDataItem", "pageSize", "pageToken", "savedQuery", )) # verify required fields with non-default values are left alone assert "dataset" in jsonified_request @@ -11402,13 +11178,6 @@ def test_search_data_items_rest_required_fields(request_type=dataset_service.Sea assert sorted(expected_params) == sorted(actual_params) -def test_search_data_items_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_data_items._get_unset_required_fields({}) - assert set(unset_fields) == (set(("annotationFilters", "annotationsFilter", "annotationsLimit", "dataItemFilter", "dataLabelingJob", "fieldMask", "orderBy", "orderByAnnotation", "orderByDataItem", "pageSize", "pageToken", "savedQuery", )) & set(("dataset", ))) - - def test_search_data_items_rest_pager(transport: str = 'rest'): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11522,17 +11291,20 @@ def test_list_saved_queries_rest_required_fields(request_type=dataset_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_saved_queries._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSavedQueries, + "_BaseListSavedQueries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_saved_queries._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11582,13 +11354,6 @@ def test_list_saved_queries_rest_required_fields(request_type=dataset_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_saved_queries_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_saved_queries._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_saved_queries_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11760,16 +11525,18 @@ def test_delete_saved_query_rest_required_fields(request_type=dataset_service.De # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_saved_query._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSavedQuery, + "_BaseDeleteSavedQuery__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_saved_query._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11815,13 +11582,6 @@ def test_delete_saved_query_rest_required_fields(request_type=dataset_service.De assert sorted(expected_params) == sorted(actual_params) -def test_delete_saved_query_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_saved_query._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_saved_query_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11922,17 +11682,20 @@ def test_get_annotation_spec_rest_required_fields(request_type=dataset_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_annotation_spec._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetAnnotationSpec, + "_BaseGetAnnotationSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_annotation_spec._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -11982,13 +11745,6 @@ def test_get_annotation_spec_rest_required_fields(request_type=dataset_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_annotation_spec_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_annotation_spec._get_unset_required_fields({}) - assert set(unset_fields) == (set(("readMask", )) & set(("name", ))) - - def test_get_annotation_spec_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12091,17 +11847,20 @@ def test_list_annotations_rest_required_fields(request_type=dataset_service.List # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_annotations._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListAnnotations, + "_BaseListAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_annotations._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12151,13 +11910,6 @@ def test_list_annotations_rest_required_fields(request_type=dataset_service.List assert sorted(expected_params) == sorted(actual_params) -def test_list_annotations_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_annotations._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_annotations_rest_flattened(): client = DatasetServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12329,16 +12081,18 @@ def test_assess_data_rest_required_fields(request_type=dataset_service.AssessDat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assess_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAssessData, + "_BaseAssessData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assess_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -12385,13 +12139,6 @@ def test_assess_data_rest_required_fields(request_type=dataset_service.AssessDat assert sorted(expected_params) == sorted(actual_params) -def test_assess_data_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.assess_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_assemble_data_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -12444,16 +12191,18 @@ def test_assemble_data_rest_required_fields(request_type=dataset_service.Assembl # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assemble_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAssembleData, + "_BaseAssembleData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assemble_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -12500,13 +12249,6 @@ def test_assemble_data_rest_required_fields(request_type=dataset_service.Assembl assert sorted(expected_params) == sorted(actual_params) -def test_assemble_data_rest_unset_required_fields(): - transport = transports.DatasetServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.assemble_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.DatasetServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_deployment_resource_pool_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_deployment_resource_pool_service.py index d2b0fbfd2d..4ed03b72b5 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_deployment_resource_pool_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_deployment_resource_pool_service.py @@ -95,6 +95,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -138,150 +150,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(None) is None - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert DeploymentResourcePoolServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - DeploymentResourcePoolServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert DeploymentResourcePoolServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - DeploymentResourcePoolServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert DeploymentResourcePoolServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - DeploymentResourcePoolServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert DeploymentResourcePoolServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -295,40 +163,6 @@ def test__get_client_cert_source(): assert DeploymentResourcePoolServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert DeploymentResourcePoolServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(DeploymentResourcePoolServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DeploymentResourcePoolServiceClient)) -@mock.patch.object(DeploymentResourcePoolServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(DeploymentResourcePoolServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - default_endpoint = DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = DeploymentResourcePoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert DeploymentResourcePoolServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, default_universe, "always") == DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == DeploymentResourcePoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert DeploymentResourcePoolServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - DeploymentResourcePoolServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert DeploymentResourcePoolServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert DeploymentResourcePoolServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert DeploymentResourcePoolServiceClient._get_universe_domain(None, None) == DeploymentResourcePoolServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - DeploymentResourcePoolServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -713,11 +547,12 @@ def test_deployment_resource_pool_service_client_get_mtls_endpoint_and_cert_sour for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -760,11 +595,12 @@ def test_deployment_resource_pool_service_client_get_mtls_endpoint_and_cert_sour for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3385,7 +3221,12 @@ def test_create_deployment_resource_pool_rest_required_fields(request_type=deplo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDeploymentResourcePool, + "_BaseCreateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -3393,9 +3234,6 @@ def test_create_deployment_resource_pool_rest_required_fields(request_type=deplo jsonified_request["parent"] = 'parent_value' jsonified_request["deploymentResourcePoolId"] = 'deployment_resource_pool_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_deployment_resource_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3444,13 +3282,6 @@ def test_create_deployment_resource_pool_rest_required_fields(request_type=deplo assert sorted(expected_params) == sorted(actual_params) -def test_create_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "deploymentResourcePool", "deploymentResourcePoolId", ))) - - def test_create_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3555,16 +3386,18 @@ def test_get_deployment_resource_pool_rest_required_fields(request_type=deployme # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDeploymentResourcePool, + "_BaseGetDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_deployment_resource_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3613,13 +3446,6 @@ def test_get_deployment_resource_pool_rest_required_fields(request_type=deployme assert sorted(expected_params) == sorted(actual_params) -def test_get_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3722,17 +3548,20 @@ def test_list_deployment_resource_pools_rest_required_fields(request_type=deploy # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_deployment_resource_pools._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDeploymentResourcePools, + "_BaseListDeploymentResourcePools__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_deployment_resource_pools._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3782,13 +3611,6 @@ def test_list_deployment_resource_pools_rest_required_fields(request_type=deploy assert sorted(expected_params) == sorted(actual_params) -def test_list_deployment_resource_pools_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_deployment_resource_pools._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_deployment_resource_pools_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3959,15 +3781,18 @@ def test_update_deployment_resource_pool_rest_required_fields(request_type=deplo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateDeploymentResourcePool, + "_BaseUpdateDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_deployment_resource_pool._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4013,13 +3838,6 @@ def test_update_deployment_resource_pool_rest_required_fields(request_type=deplo assert sorted(expected_params) == sorted(actual_params) -def test_update_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("deploymentResourcePool", "updateMask", ))) - - def test_update_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4126,16 +3944,18 @@ def test_delete_deployment_resource_pool_rest_required_fields(request_type=deplo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_deployment_resource_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDeploymentResourcePool, + "_BaseDeleteDeploymentResourcePool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_deployment_resource_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4181,13 +4001,6 @@ def test_delete_deployment_resource_pool_rest_required_fields(request_type=deplo assert sorted(expected_params) == sorted(actual_params) -def test_delete_deployment_resource_pool_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_deployment_resource_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_deployment_resource_pool_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4288,17 +4101,20 @@ def test_query_deployed_models_rest_required_fields(request_type=deployment_reso # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_deployed_models._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryDeployedModels, + "_BaseQueryDeployedModels__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["deploymentResourcePool"] = 'deployment_resource_pool_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_deployed_models._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "deploymentResourcePool" in jsonified_request @@ -4348,13 +4164,6 @@ def test_query_deployed_models_rest_required_fields(request_type=deployment_reso assert sorted(expected_params) == sorted(actual_params) -def test_query_deployed_models_rest_unset_required_fields(): - transport = transports.DeploymentResourcePoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_deployed_models._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("deploymentResourcePool", ))) - - def test_query_deployed_models_rest_flattened(): client = DeploymentResourcePoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_endpoint_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_endpoint_service.py index 22c1526511..298bef6f87 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_endpoint_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_endpoint_service.py @@ -99,6 +99,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -142,150 +154,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert EndpointServiceClient._get_default_mtls_endpoint(None) is None - assert EndpointServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert EndpointServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert EndpointServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EndpointServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - EndpointServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert EndpointServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert EndpointServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert EndpointServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - EndpointServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert EndpointServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert EndpointServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - EndpointServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert EndpointServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert EndpointServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -299,40 +167,6 @@ def test__get_client_cert_source(): assert EndpointServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert EndpointServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(EndpointServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EndpointServiceClient)) -@mock.patch.object(EndpointServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EndpointServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = EndpointServiceClient._DEFAULT_UNIVERSE - default_endpoint = EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = EndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert EndpointServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert EndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert EndpointServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert EndpointServiceClient._get_api_endpoint(None, None, default_universe, "always") == EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert EndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == EndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert EndpointServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert EndpointServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - EndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert EndpointServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert EndpointServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert EndpointServiceClient._get_universe_domain(None, None) == EndpointServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - EndpointServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -717,11 +551,12 @@ def test_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -764,11 +599,12 @@ def test_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -4899,17 +4735,20 @@ def test_create_endpoint_rest_required_fields(request_type=endpoint_service.Crea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateEndpoint, + "_BaseCreateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_endpoint._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("endpoint_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("endpointId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4957,13 +4796,6 @@ def test_create_endpoint_rest_required_fields(request_type=endpoint_service.Crea assert sorted(expected_params) == sorted(actual_params) -def test_create_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(("endpointId", )) & set(("parent", "endpoint", ))) - - def test_create_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5068,16 +4900,18 @@ def test_get_endpoint_rest_required_fields(request_type=endpoint_service.GetEndp # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetEndpoint, + "_BaseGetEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5126,13 +4960,6 @@ def test_get_endpoint_rest_required_fields(request_type=endpoint_service.GetEndp assert sorted(expected_params) == sorted(actual_params) -def test_get_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5235,17 +5062,20 @@ def test_list_endpoints_rest_required_fields(request_type=endpoint_service.ListE # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_endpoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListEndpoints, + "_BaseListEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_endpoints._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5295,13 +5125,6 @@ def test_list_endpoints_rest_required_fields(request_type=endpoint_service.ListE assert sorted(expected_params) == sorted(actual_params) -def test_list_endpoints_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_endpoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_endpoints_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5468,15 +5291,18 @@ def test_update_endpoint_rest_required_fields(request_type=endpoint_service.Upda # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateEndpoint, + "_BaseUpdateEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -5525,13 +5351,6 @@ def test_update_endpoint_rest_required_fields(request_type=endpoint_service.Upda assert sorted(expected_params) == sorted(actual_params) -def test_update_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("endpoint", "updateMask", ))) - - def test_update_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5639,14 +5458,16 @@ def test_update_endpoint_long_running_rest_required_fields(request_type=endpoint # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint_long_running._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateEndpointLongRunning, + "_BaseUpdateEndpointLongRunning__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_endpoint_long_running._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone client = EndpointServiceClient( @@ -5691,13 +5512,6 @@ def test_update_endpoint_long_running_rest_required_fields(request_type=endpoint assert sorted(expected_params) == sorted(actual_params) -def test_update_endpoint_long_running_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_endpoint_long_running._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_update_endpoint_long_running_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5802,16 +5616,18 @@ def test_delete_endpoint_rest_required_fields(request_type=endpoint_service.Dele # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteEndpoint, + "_BaseDeleteEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5857,13 +5673,6 @@ def test_delete_endpoint_rest_required_fields(request_type=endpoint_service.Dele assert sorted(expected_params) == sorted(actual_params) -def test_delete_endpoint_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_endpoint_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5968,16 +5777,18 @@ def test_deploy_model_rest_required_fields(request_type=endpoint_service.DeployM # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeployModel, + "_BaseDeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -6024,13 +5835,6 @@ def test_deploy_model_rest_required_fields(request_type=endpoint_service.DeployM assert sorted(expected_params) == sorted(actual_params) -def test_deploy_model_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.deploy_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "deployedModel", ))) - - def test_deploy_model_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6140,7 +5944,12 @@ def test_undeploy_model_rest_required_fields(request_type=endpoint_service.Undep # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUndeployModel, + "_BaseUndeployModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -6148,9 +5957,6 @@ def test_undeploy_model_rest_required_fields(request_type=endpoint_service.Undep jsonified_request["endpoint"] = 'endpoint_value' jsonified_request["deployedModelId"] = 'deployed_model_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -6199,13 +6005,6 @@ def test_undeploy_model_rest_required_fields(request_type=endpoint_service.Undep assert sorted(expected_params) == sorted(actual_params) -def test_undeploy_model_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.undeploy_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "deployedModelId", ))) - - def test_undeploy_model_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6314,16 +6113,18 @@ def test_mutate_deployed_model_rest_required_fields(request_type=endpoint_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseMutateDeployedModel, + "_BaseMutateDeployedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -6370,13 +6171,6 @@ def test_mutate_deployed_model_rest_required_fields(request_type=endpoint_servic assert sorted(expected_params) == sorted(actual_params) -def test_mutate_deployed_model_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.mutate_deployed_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "deployedModel", "updateMask", ))) - - def test_mutate_deployed_model_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6485,16 +6279,18 @@ def test_set_publisher_model_config_rest_required_fields(request_type=endpoint_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).set_publisher_model_config._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSetPublisherModelConfig, + "_BaseSetPublisherModelConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).set_publisher_model_config._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6541,13 +6337,6 @@ def test_set_publisher_model_config_rest_required_fields(request_type=endpoint_s assert sorted(expected_params) == sorted(actual_params) -def test_set_publisher_model_config_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.set_publisher_model_config._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "publisherModelConfig", ))) - - def test_set_publisher_model_config_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6650,16 +6439,18 @@ def test_fetch_publisher_model_config_rest_required_fields(request_type=endpoint # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_publisher_model_config._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseFetchPublisherModelConfig, + "_BaseFetchPublisherModelConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_publisher_model_config._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6708,13 +6499,6 @@ def test_fetch_publisher_model_config_rest_required_fields(request_type=endpoint assert sorted(expected_params) == sorted(actual_params) -def test_fetch_publisher_model_config_rest_unset_required_fields(): - transport = transports.EndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.fetch_publisher_model_config._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_fetch_publisher_model_config_rest_flattened(): client = EndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_evaluation_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_evaluation_service.py index c191e63317..674fce51e7 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_evaluation_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_evaluation_service.py @@ -91,6 +91,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -134,150 +146,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert EvaluationServiceClient._get_default_mtls_endpoint(None) is None - assert EvaluationServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert EvaluationServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert EvaluationServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EvaluationServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - EvaluationServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert EvaluationServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - EvaluationServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert EvaluationServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert EvaluationServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - EvaluationServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert EvaluationServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert EvaluationServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -291,40 +159,6 @@ def test__get_client_cert_source(): assert EvaluationServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert EvaluationServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(EvaluationServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EvaluationServiceClient)) -@mock.patch.object(EvaluationServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(EvaluationServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = EvaluationServiceClient._DEFAULT_UNIVERSE - default_endpoint = EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = EvaluationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert EvaluationServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert EvaluationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - assert EvaluationServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert EvaluationServiceClient._get_api_endpoint(None, None, default_universe, "always") == EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - assert EvaluationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == EvaluationServiceClient.DEFAULT_MTLS_ENDPOINT - assert EvaluationServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert EvaluationServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - EvaluationServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert EvaluationServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert EvaluationServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert EvaluationServiceClient._get_universe_domain(None, None) == EvaluationServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - EvaluationServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -709,11 +543,12 @@ def test_evaluation_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -756,11 +591,12 @@ def test_evaluation_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1720,16 +1556,18 @@ def test_evaluate_instances_rest_required_fields(request_type=evaluation_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).evaluate_instances._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseEvaluateInstances, + "_BaseEvaluateInstances__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["location"] = 'location_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).evaluate_instances._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "location" in jsonified_request assert jsonified_request["location"] == 'location_value' @@ -1779,13 +1617,6 @@ def test_evaluate_instances_rest_required_fields(request_type=evaluation_service assert sorted(expected_params) == sorted(actual_params) -def test_evaluate_instances_rest_unset_required_fields(): - transport = transports.EvaluationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.evaluate_instances._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("location", ))) - - def test_evaluate_dataset_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -1838,16 +1669,18 @@ def test_evaluate_dataset_rest_required_fields(request_type=evaluation_service.E # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).evaluate_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseEvaluateDataset, + "_BaseEvaluateDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["location"] = 'location_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).evaluate_dataset._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "location" in jsonified_request assert jsonified_request["location"] == 'location_value' @@ -1894,13 +1727,6 @@ def test_evaluate_dataset_rest_required_fields(request_type=evaluation_service.E assert sorted(expected_params) == sorted(actual_params) -def test_evaluate_dataset_rest_unset_required_fields(): - transport = transports.EvaluationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.evaluate_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("location", "dataset", "metrics", "outputConfig", ))) - - def test_generate_instance_rubrics_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -1949,16 +1775,18 @@ def test_generate_instance_rubrics_rest_required_fields(request_type=evaluation_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_instance_rubrics._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGenerateInstanceRubrics, + "_BaseGenerateInstanceRubrics__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["location"] = 'location_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_instance_rubrics._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "location" in jsonified_request assert jsonified_request["location"] == 'location_value' @@ -2008,13 +1836,6 @@ def test_generate_instance_rubrics_rest_required_fields(request_type=evaluation_ assert sorted(expected_params) == sorted(actual_params) -def test_generate_instance_rubrics_rest_unset_required_fields(): - transport = transports.EvaluationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.generate_instance_rubrics._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("location", "contents", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.EvaluationServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_example_store_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_example_store_service.py index dc46e19d06..c1631ebb39 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_example_store_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_example_store_service.py @@ -93,6 +93,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -136,150 +148,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ExampleStoreServiceClient._get_default_mtls_endpoint(None) is None - assert ExampleStoreServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ExampleStoreServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ExampleStoreServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ExampleStoreServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ExampleStoreServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ExampleStoreServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ExampleStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ExampleStoreServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ExampleStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ExampleStoreServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ExampleStoreServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ExampleStoreServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ExampleStoreServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ExampleStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ExampleStoreServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ExampleStoreServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ExampleStoreServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ExampleStoreServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ExampleStoreServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ExampleStoreServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ExampleStoreServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ExampleStoreServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ExampleStoreServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ExampleStoreServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ExampleStoreServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ExampleStoreServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ExampleStoreServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ExampleStoreServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -293,40 +161,6 @@ def test__get_client_cert_source(): assert ExampleStoreServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ExampleStoreServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ExampleStoreServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ExampleStoreServiceClient)) -@mock.patch.object(ExampleStoreServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ExampleStoreServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ExampleStoreServiceClient._DEFAULT_UNIVERSE - default_endpoint = ExampleStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ExampleStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ExampleStoreServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ExampleStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ExampleStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExampleStoreServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ExampleStoreServiceClient._get_api_endpoint(None, None, default_universe, "always") == ExampleStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExampleStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ExampleStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExampleStoreServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ExampleStoreServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ExampleStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ExampleStoreServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ExampleStoreServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ExampleStoreServiceClient._get_universe_domain(None, None) == ExampleStoreServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ExampleStoreServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -711,11 +545,12 @@ def test_example_store_service_client_get_mtls_endpoint_and_cert_source(client_c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -758,11 +593,12 @@ def test_example_store_service_client_get_mtls_endpoint_and_cert_source(client_c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3947,16 +3783,18 @@ def test_create_example_store_rest_required_fields(request_type=example_store_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_example_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateExampleStore, + "_BaseCreateExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_example_store._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -4003,13 +3841,6 @@ def test_create_example_store_rest_required_fields(request_type=example_store_se assert sorted(expected_params) == sorted(actual_params) -def test_create_example_store_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_example_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "exampleStore", ))) - - def test_create_example_store_rest_flattened(): client = ExampleStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4112,16 +3943,18 @@ def test_get_example_store_rest_required_fields(request_type=example_store_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_example_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetExampleStore, + "_BaseGetExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_example_store._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4170,13 +4003,6 @@ def test_get_example_store_rest_required_fields(request_type=example_store_servi assert sorted(expected_params) == sorted(actual_params) -def test_get_example_store_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_example_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_example_store_rest_flattened(): client = ExampleStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4282,15 +4108,18 @@ def test_update_example_store_rest_required_fields(request_type=example_store_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_example_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateExampleStore, + "_BaseUpdateExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_example_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4336,13 +4165,6 @@ def test_update_example_store_rest_required_fields(request_type=example_store_se assert sorted(expected_params) == sorted(actual_params) -def test_update_example_store_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_example_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("exampleStore", ))) - - def test_update_example_store_rest_flattened(): client = ExampleStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4449,16 +4271,18 @@ def test_delete_example_store_rest_required_fields(request_type=example_store_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_example_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteExampleStore, + "_BaseDeleteExampleStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_example_store._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4504,13 +4328,6 @@ def test_delete_example_store_rest_required_fields(request_type=example_store_se assert sorted(expected_params) == sorted(actual_params) -def test_delete_example_store_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_example_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_example_store_rest_flattened(): client = ExampleStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4611,17 +4428,20 @@ def test_list_example_stores_rest_required_fields(request_type=example_store_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_example_stores._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListExampleStores, + "_BaseListExampleStores__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_example_stores._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4671,13 +4491,6 @@ def test_list_example_stores_rest_required_fields(request_type=example_store_ser assert sorted(expected_params) == sorted(actual_params) -def test_list_example_stores_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_example_stores._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_example_stores_rest_flattened(): client = ExampleStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4845,16 +4658,18 @@ def test_upsert_examples_rest_required_fields(request_type=example_store_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upsert_examples._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpsertExamples, + "_BaseUpsertExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["exampleStore"] = 'example_store_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upsert_examples._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "exampleStore" in jsonified_request assert jsonified_request["exampleStore"] == 'example_store_value' @@ -4904,13 +4719,6 @@ def test_upsert_examples_rest_required_fields(request_type=example_store_service assert sorted(expected_params) == sorted(actual_params) -def test_upsert_examples_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upsert_examples._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("exampleStore", "examples", ))) - - def test_remove_examples_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -4959,16 +4767,18 @@ def test_remove_examples_rest_required_fields(request_type=example_store_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_examples._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRemoveExamples, + "_BaseRemoveExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["exampleStore"] = 'example_store_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_examples._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "exampleStore" in jsonified_request assert jsonified_request["exampleStore"] == 'example_store_value' @@ -5018,13 +4828,6 @@ def test_remove_examples_rest_required_fields(request_type=example_store_service assert sorted(expected_params) == sorted(actual_params) -def test_remove_examples_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.remove_examples._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("exampleStore", ))) - - def test_search_examples_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -5073,16 +4876,18 @@ def test_search_examples_rest_required_fields(request_type=example_store_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_examples._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchExamples, + "_BaseSearchExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["exampleStore"] = 'example_store_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_examples._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "exampleStore" in jsonified_request assert jsonified_request["exampleStore"] == 'example_store_value' @@ -5132,13 +4937,6 @@ def test_search_examples_rest_required_fields(request_type=example_store_service assert sorted(expected_params) == sorted(actual_params) -def test_search_examples_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_examples._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("exampleStore", ))) - - def test_fetch_examples_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -5187,16 +4985,18 @@ def test_fetch_examples_rest_required_fields(request_type=example_store_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_examples._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseFetchExamples, + "_BaseFetchExamples__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["exampleStore"] = 'example_store_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_examples._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "exampleStore" in jsonified_request assert jsonified_request["exampleStore"] == 'example_store_value' @@ -5246,13 +5046,6 @@ def test_fetch_examples_rest_required_fields(request_type=example_store_service. assert sorted(expected_params) == sorted(actual_params) -def test_fetch_examples_rest_unset_required_fields(): - transport = transports.ExampleStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.fetch_examples._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("exampleStore", ))) - - def test_fetch_examples_rest_pager(transport: str = 'rest'): client = ExampleStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_extension_execution_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_extension_execution_service.py index 3a9c73afbc..ef74da1a9d 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_extension_execution_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_extension_execution_service.py @@ -83,6 +83,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -126,150 +138,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ExtensionExecutionServiceClient._get_default_mtls_endpoint(None) is None - assert ExtensionExecutionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ExtensionExecutionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ExtensionExecutionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ExtensionExecutionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ExtensionExecutionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ExtensionExecutionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ExtensionExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ExtensionExecutionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ExtensionExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ExtensionExecutionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ExtensionExecutionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ExtensionExecutionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ExtensionExecutionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ExtensionExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ExtensionExecutionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ExtensionExecutionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ExtensionExecutionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ExtensionExecutionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -283,40 +151,6 @@ def test__get_client_cert_source(): assert ExtensionExecutionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ExtensionExecutionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ExtensionExecutionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ExtensionExecutionServiceClient)) -@mock.patch.object(ExtensionExecutionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ExtensionExecutionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ExtensionExecutionServiceClient._DEFAULT_UNIVERSE - default_endpoint = ExtensionExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ExtensionExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ExtensionExecutionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ExtensionExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ExtensionExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExtensionExecutionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ExtensionExecutionServiceClient._get_api_endpoint(None, None, default_universe, "always") == ExtensionExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExtensionExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ExtensionExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExtensionExecutionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ExtensionExecutionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ExtensionExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ExtensionExecutionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ExtensionExecutionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ExtensionExecutionServiceClient._get_universe_domain(None, None) == ExtensionExecutionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ExtensionExecutionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -701,11 +535,12 @@ def test_extension_execution_service_client_get_mtls_endpoint_and_cert_source(cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -748,11 +583,12 @@ def test_extension_execution_service_client_get_mtls_endpoint_and_cert_source(cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1674,7 +1510,12 @@ def test_execute_extension_rest_required_fields(request_type=extension_execution # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).execute_extension._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExecuteExtension, + "_BaseExecuteExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -1682,9 +1523,6 @@ def test_execute_extension_rest_required_fields(request_type=extension_execution jsonified_request["name"] = 'name_value' jsonified_request["operationId"] = 'operation_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).execute_extension._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -1736,13 +1574,6 @@ def test_execute_extension_rest_required_fields(request_type=extension_execution assert sorted(expected_params) == sorted(actual_params) -def test_execute_extension_rest_unset_required_fields(): - transport = transports.ExtensionExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.execute_extension._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "operationId", ))) - - def test_execute_extension_rest_flattened(): client = ExtensionExecutionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -1847,16 +1678,18 @@ def test_query_extension_rest_required_fields(request_type=extension_execution_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_extension._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryExtension, + "_BaseQueryExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_extension._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -1906,13 +1739,6 @@ def test_query_extension_rest_required_fields(request_type=extension_execution_s assert sorted(expected_params) == sorted(actual_params) -def test_query_extension_rest_unset_required_fields(): - transport = transports.ExtensionExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_extension._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "contents", ))) - - def test_query_extension_rest_flattened(): client = ExtensionExecutionServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_extension_registry_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_extension_registry_service.py index 3eb35262ba..5859bad15f 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_extension_registry_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_extension_registry_service.py @@ -92,6 +92,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -135,150 +147,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ExtensionRegistryServiceClient._get_default_mtls_endpoint(None) is None - assert ExtensionRegistryServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ExtensionRegistryServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ExtensionRegistryServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ExtensionRegistryServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ExtensionRegistryServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ExtensionRegistryServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ExtensionRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ExtensionRegistryServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ExtensionRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ExtensionRegistryServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ExtensionRegistryServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ExtensionRegistryServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ExtensionRegistryServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ExtensionRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ExtensionRegistryServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ExtensionRegistryServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ExtensionRegistryServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ExtensionRegistryServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -292,40 +160,6 @@ def test__get_client_cert_source(): assert ExtensionRegistryServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ExtensionRegistryServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ExtensionRegistryServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ExtensionRegistryServiceClient)) -@mock.patch.object(ExtensionRegistryServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ExtensionRegistryServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ExtensionRegistryServiceClient._DEFAULT_UNIVERSE - default_endpoint = ExtensionRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ExtensionRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ExtensionRegistryServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ExtensionRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ExtensionRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExtensionRegistryServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ExtensionRegistryServiceClient._get_api_endpoint(None, None, default_universe, "always") == ExtensionRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExtensionRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ExtensionRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert ExtensionRegistryServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ExtensionRegistryServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ExtensionRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ExtensionRegistryServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ExtensionRegistryServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ExtensionRegistryServiceClient._get_universe_domain(None, None) == ExtensionRegistryServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ExtensionRegistryServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -710,11 +544,12 @@ def test_extension_registry_service_client_get_mtls_endpoint_and_cert_source(cli for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -757,11 +592,12 @@ def test_extension_registry_service_client_get_mtls_endpoint_and_cert_source(cli for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2854,16 +2690,18 @@ def test_import_extension_rest_required_fields(request_type=extension_registry_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_extension._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportExtension, + "_BaseImportExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_extension._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2910,13 +2748,6 @@ def test_import_extension_rest_required_fields(request_type=extension_registry_s assert sorted(expected_params) == sorted(actual_params) -def test_import_extension_rest_unset_required_fields(): - transport = transports.ExtensionRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_extension._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "extension", ))) - - def test_import_extension_rest_flattened(): client = ExtensionRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3019,16 +2850,18 @@ def test_get_extension_rest_required_fields(request_type=extension_registry_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_extension._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetExtension, + "_BaseGetExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_extension._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3077,13 +2910,6 @@ def test_get_extension_rest_required_fields(request_type=extension_registry_serv assert sorted(expected_params) == sorted(actual_params) -def test_get_extension_rest_unset_required_fields(): - transport = transports.ExtensionRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_extension._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_extension_rest_flattened(): client = ExtensionRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3186,17 +3012,20 @@ def test_list_extensions_rest_required_fields(request_type=extension_registry_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_extensions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListExtensions, + "_BaseListExtensions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_extensions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3246,13 +3075,6 @@ def test_list_extensions_rest_required_fields(request_type=extension_registry_se assert sorted(expected_params) == sorted(actual_params) -def test_list_extensions_rest_unset_required_fields(): - transport = transports.ExtensionRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_extensions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_extensions_rest_flattened(): client = ExtensionRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3419,15 +3241,18 @@ def test_update_extension_rest_required_fields(request_type=extension_registry_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_extension._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateExtension, + "_BaseUpdateExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_extension._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3476,13 +3301,6 @@ def test_update_extension_rest_required_fields(request_type=extension_registry_s assert sorted(expected_params) == sorted(actual_params) -def test_update_extension_rest_unset_required_fields(): - transport = transports.ExtensionRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_extension._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("extension", "updateMask", ))) - - def test_update_extension_rest_flattened(): client = ExtensionRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3591,16 +3409,18 @@ def test_delete_extension_rest_required_fields(request_type=extension_registry_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_extension._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteExtension, + "_BaseDeleteExtension__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_extension._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3646,13 +3466,6 @@ def test_delete_extension_rest_required_fields(request_type=extension_registry_s assert sorted(expected_params) == sorted(actual_params) -def test_delete_extension_rest_unset_required_fields(): - transport = transports.ExtensionRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_extension._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_extension_rest_flattened(): client = ExtensionRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_admin_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_admin_service.py index 5daafe2aac..9ca2a9933b 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_admin_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_admin_service.py @@ -97,6 +97,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -140,150 +152,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(None) is None - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeatureOnlineStoreAdminServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreAdminServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreAdminServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeatureOnlineStoreAdminServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeatureOnlineStoreAdminServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -297,40 +165,6 @@ def test__get_client_cert_source(): assert FeatureOnlineStoreAdminServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeatureOnlineStoreAdminServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeatureOnlineStoreAdminServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreAdminServiceClient)) -@mock.patch.object(FeatureOnlineStoreAdminServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreAdminServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeatureOnlineStoreAdminServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeatureOnlineStoreAdminServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreAdminServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeatureOnlineStoreAdminServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeatureOnlineStoreAdminServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeatureOnlineStoreAdminServiceClient._get_universe_domain(None, None) == FeatureOnlineStoreAdminServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreAdminServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -715,11 +549,12 @@ def test_feature_online_store_admin_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -762,11 +597,12 @@ def test_feature_online_store_admin_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5848,7 +5684,12 @@ def test_create_feature_online_store_rest_required_fields(request_type=feature_o # verify fields with default values are dropped assert "featureOnlineStoreId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureOnlineStore, + "_BaseCreateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -5858,10 +5699,8 @@ def test_create_feature_online_store_rest_required_fields(request_type=feature_o jsonified_request["parent"] = 'parent_value' jsonified_request["featureOnlineStoreId"] = 'feature_online_store_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_online_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_online_store_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureOnlineStoreId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5915,13 +5754,6 @@ def test_create_feature_online_store_rest_required_fields(request_type=feature_o assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureOnlineStoreId", )) & set(("parent", "featureOnlineStore", "featureOnlineStoreId", ))) - - def test_create_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6026,16 +5858,18 @@ def test_get_feature_online_store_rest_required_fields(request_type=feature_onli # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureOnlineStore, + "_BaseGetFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_online_store._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6084,13 +5918,6 @@ def test_get_feature_online_store_rest_required_fields(request_type=feature_onli assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6193,17 +6020,20 @@ def test_list_feature_online_stores_rest_required_fields(request_type=feature_on # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_online_stores._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureOnlineStores, + "_BaseListFeatureOnlineStores__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_online_stores._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6253,13 +6083,6 @@ def test_list_feature_online_stores_rest_required_fields(request_type=feature_on assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_online_stores_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_online_stores._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_online_stores_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6430,15 +6253,18 @@ def test_update_feature_online_store_rest_required_fields(request_type=feature_o # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeatureOnlineStore, + "_BaseUpdateFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_online_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -6484,13 +6310,6 @@ def test_update_feature_online_store_rest_required_fields(request_type=feature_o assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featureOnlineStore", ))) - - def test_update_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6597,17 +6416,20 @@ def test_delete_feature_online_store_rest_required_fields(request_type=feature_o # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_online_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureOnlineStore, + "_BaseDeleteFeatureOnlineStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_online_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -6654,13 +6476,6 @@ def test_delete_feature_online_store_rest_required_fields(request_type=feature_o assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_online_store_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_online_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_feature_online_store_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6769,7 +6584,12 @@ def test_create_feature_view_rest_required_fields(request_type=feature_online_st # verify fields with default values are dropped assert "featureViewId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureView, + "_BaseCreateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -6779,10 +6599,8 @@ def test_create_feature_view_rest_required_fields(request_type=feature_online_st jsonified_request["parent"] = 'parent_value' jsonified_request["featureViewId"] = 'feature_view_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_view._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_view_id", "run_sync_immediately", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureViewId", "runSyncImmediately", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6836,13 +6654,6 @@ def test_create_feature_view_rest_required_fields(request_type=feature_online_st assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureViewId", "runSyncImmediately", )) & set(("parent", "featureView", "featureViewId", ))) - - def test_create_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6947,16 +6758,18 @@ def test_get_feature_view_rest_required_fields(request_type=feature_online_store # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureView, + "_BaseGetFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7005,13 +6818,6 @@ def test_get_feature_view_rest_required_fields(request_type=feature_online_store assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7114,17 +6920,20 @@ def test_list_feature_views_rest_required_fields(request_type=feature_online_sto # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_views._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureViews, + "_BaseListFeatureViews__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_views._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -7174,13 +6983,6 @@ def test_list_feature_views_rest_required_fields(request_type=feature_online_sto assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_views_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_views._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_views_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7351,15 +7153,18 @@ def test_update_feature_view_rest_required_fields(request_type=feature_online_st # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeatureView, + "_BaseUpdateFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_view._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -7405,13 +7210,6 @@ def test_update_feature_view_rest_required_fields(request_type=feature_online_st assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featureView", ))) - - def test_update_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7518,16 +7316,18 @@ def test_delete_feature_view_rest_required_fields(request_type=feature_online_st # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureView, + "_BaseDeleteFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_view._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7573,13 +7373,6 @@ def test_delete_feature_view_rest_required_fields(request_type=feature_online_st assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7680,16 +7473,18 @@ def test_sync_feature_view_rest_required_fields(request_type=feature_online_stor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).sync_feature_view._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSyncFeatureView, + "_BaseSyncFeatureView__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featureView"] = 'feature_view_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).sync_feature_view._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featureView" in jsonified_request assert jsonified_request["featureView"] == 'feature_view_value' @@ -7739,13 +7534,6 @@ def test_sync_feature_view_rest_required_fields(request_type=feature_online_stor assert sorted(expected_params) == sorted(actual_params) -def test_sync_feature_view_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.sync_feature_view._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featureView", ))) - - def test_sync_feature_view_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7848,16 +7636,18 @@ def test_get_feature_view_sync_rest_required_fields(request_type=feature_online_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view_sync._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureViewSync, + "_BaseGetFeatureViewSync__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_view_sync._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7906,13 +7696,6 @@ def test_get_feature_view_sync_rest_required_fields(request_type=feature_online_ assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_view_sync_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_view_sync._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_view_sync_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8015,17 +7798,20 @@ def test_list_feature_view_syncs_rest_required_fields(request_type=feature_onlin # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_view_syncs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureViewSyncs, + "_BaseListFeatureViewSyncs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_view_syncs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8075,13 +7861,6 @@ def test_list_feature_view_syncs_rest_required_fields(request_type=feature_onlin assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_view_syncs_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreAdminServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_view_syncs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_view_syncs_rest_flattened(): client = FeatureOnlineStoreAdminServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_service.py index e0db1e457c..b1a4d8e42d 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_feature_online_store_service.py @@ -83,6 +83,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -126,150 +138,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(None) is None - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeatureOnlineStoreServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeatureOnlineStoreServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeatureOnlineStoreServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeatureOnlineStoreServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeatureOnlineStoreServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -283,40 +151,6 @@ def test__get_client_cert_source(): assert FeatureOnlineStoreServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeatureOnlineStoreServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeatureOnlineStoreServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreServiceClient)) -@mock.patch.object(FeatureOnlineStoreServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureOnlineStoreServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeatureOnlineStoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeatureOnlineStoreServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeatureOnlineStoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeatureOnlineStoreServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureOnlineStoreServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeatureOnlineStoreServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeatureOnlineStoreServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeatureOnlineStoreServiceClient._get_universe_domain(None, None) == FeatureOnlineStoreServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeatureOnlineStoreServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -701,11 +535,12 @@ def test_feature_online_store_service_client_get_mtls_endpoint_and_cert_source(c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -748,11 +583,12 @@ def test_feature_online_store_service_client_get_mtls_endpoint_and_cert_source(c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2066,16 +1902,18 @@ def test_fetch_feature_values_rest_required_fields(request_type=feature_online_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseFetchFeatureValues, + "_BaseFetchFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featureView"] = 'feature_view_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).fetch_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featureView" in jsonified_request assert jsonified_request["featureView"] == 'feature_view_value' @@ -2125,13 +1963,6 @@ def test_fetch_feature_values_rest_required_fields(request_type=feature_online_s assert sorted(expected_params) == sorted(actual_params) -def test_fetch_feature_values_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.fetch_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featureView", ))) - - def test_fetch_feature_values_rest_flattened(): client = FeatureOnlineStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2247,16 +2078,18 @@ def test_search_nearest_entities_rest_required_fields(request_type=feature_onlin # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_nearest_entities._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchNearestEntities, + "_BaseSearchNearestEntities__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featureView"] = 'feature_view_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_nearest_entities._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featureView" in jsonified_request assert jsonified_request["featureView"] == 'feature_view_value' @@ -2306,13 +2139,6 @@ def test_search_nearest_entities_rest_required_fields(request_type=feature_onlin assert sorted(expected_params) == sorted(actual_params) -def test_search_nearest_entities_rest_unset_required_fields(): - transport = transports.FeatureOnlineStoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_nearest_entities._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featureView", "query", ))) - - def test_feature_view_direct_write_rest_unimplemented(): client = FeatureOnlineStoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_feature_registry_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_feature_registry_service.py index 9304982e72..9662ab0cb2 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_feature_registry_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_feature_registry_service.py @@ -104,6 +104,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -147,150 +159,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(None) is None - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeatureRegistryServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeatureRegistryServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeatureRegistryServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureRegistryServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeatureRegistryServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeatureRegistryServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeatureRegistryServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -304,40 +172,6 @@ def test__get_client_cert_source(): assert FeatureRegistryServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeatureRegistryServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeatureRegistryServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureRegistryServiceClient)) -@mock.patch.object(FeatureRegistryServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeatureRegistryServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeatureRegistryServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeatureRegistryServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeatureRegistryServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeatureRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeatureRegistryServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeatureRegistryServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeatureRegistryServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeatureRegistryServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeatureRegistryServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeatureRegistryServiceClient._get_universe_domain(None, None) == FeatureRegistryServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeatureRegistryServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -722,11 +556,12 @@ def test_feature_registry_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -769,11 +604,12 @@ def test_feature_registry_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -8034,7 +7870,12 @@ def test_create_feature_group_rest_required_fields(request_type=feature_registry # verify fields with default values are dropped assert "featureGroupId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureGroup, + "_BaseCreateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -8044,10 +7885,8 @@ def test_create_feature_group_rest_required_fields(request_type=feature_registry jsonified_request["parent"] = 'parent_value' jsonified_request["featureGroupId"] = 'feature_group_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_group._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_group_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureGroupId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8101,13 +7940,6 @@ def test_create_feature_group_rest_required_fields(request_type=feature_registry assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureGroupId", )) & set(("parent", "featureGroup", "featureGroupId", ))) - - def test_create_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8212,16 +8044,18 @@ def test_get_feature_group_rest_required_fields(request_type=feature_registry_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureGroup, + "_BaseGetFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_group._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8270,13 +8104,6 @@ def test_get_feature_group_rest_required_fields(request_type=feature_registry_se assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8379,17 +8206,20 @@ def test_list_feature_groups_rest_required_fields(request_type=feature_registry_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_groups._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureGroups, + "_BaseListFeatureGroups__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_groups._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8439,13 +8269,6 @@ def test_list_feature_groups_rest_required_fields(request_type=feature_registry_ assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_groups_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_groups._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_groups_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8616,15 +8439,18 @@ def test_update_feature_group_rest_required_fields(request_type=feature_registry # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeatureGroup, + "_BaseUpdateFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_group._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -8670,13 +8496,6 @@ def test_update_feature_group_rest_required_fields(request_type=feature_registry assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featureGroup", ))) - - def test_update_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8783,17 +8602,20 @@ def test_delete_feature_group_rest_required_fields(request_type=feature_registry # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_group._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureGroup, + "_BaseDeleteFeatureGroup__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_group._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -8840,13 +8662,6 @@ def test_delete_feature_group_rest_required_fields(request_type=feature_registry assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_group_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_group._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_feature_group_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8955,7 +8770,12 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C # verify fields with default values are dropped assert "featureId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -8965,10 +8785,8 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C jsonified_request["parent"] = 'parent_value' jsonified_request["featureId"] = 'feature_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9022,13 +8840,6 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureId", )) & set(("parent", "feature", "featureId", ))) - - def test_create_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9137,16 +8948,18 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -9193,13 +9006,6 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_features_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_features_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9302,17 +9108,20 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_stats_and_anomaly_spec", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureStatsAndAnomalySpec", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9362,13 +9171,6 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureStatsAndAnomalySpec", )) & set(("name", ))) - - def test_get_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9471,17 +9273,20 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "latest_stats_count", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9531,13 +9336,6 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_features_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_features_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9708,15 +9506,18 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -9762,13 +9563,6 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("feature", ))) - - def test_update_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9875,16 +9669,18 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9930,13 +9726,6 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_feature_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10043,7 +9832,12 @@ def test_create_feature_monitor_rest_required_fields(request_type=feature_regist # verify fields with default values are dropped assert "featureMonitorId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureMonitor, + "_BaseCreateFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -10053,10 +9847,8 @@ def test_create_feature_monitor_rest_required_fields(request_type=feature_regist jsonified_request["parent"] = 'parent_value' jsonified_request["featureMonitorId"] = 'feature_monitor_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_monitor._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_monitor_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureMonitorId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10110,13 +9902,6 @@ def test_create_feature_monitor_rest_required_fields(request_type=feature_regist assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_monitor_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureMonitorId", )) & set(("parent", "featureMonitor", "featureMonitorId", ))) - - def test_create_feature_monitor_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10221,16 +10006,18 @@ def test_get_feature_monitor_rest_required_fields(request_type=feature_registry_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureMonitor, + "_BaseGetFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_monitor._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10279,13 +10066,6 @@ def test_get_feature_monitor_rest_required_fields(request_type=feature_registry_ assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_monitor_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_monitor_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10388,17 +10168,20 @@ def test_list_feature_monitors_rest_required_fields(request_type=feature_registr # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_monitors._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureMonitors, + "_BaseListFeatureMonitors__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_monitors._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10448,13 +10231,6 @@ def test_list_feature_monitors_rest_required_fields(request_type=feature_registr assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_monitors_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_monitors._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_monitors_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10625,15 +10401,18 @@ def test_update_feature_monitor_rest_required_fields(request_type=feature_regist # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeatureMonitor, + "_BaseUpdateFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature_monitor._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -10679,13 +10458,6 @@ def test_update_feature_monitor_rest_required_fields(request_type=feature_regist assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_monitor_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featureMonitor", ))) - - def test_update_feature_monitor_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10792,16 +10564,18 @@ def test_delete_feature_monitor_rest_required_fields(request_type=feature_regist # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureMonitor, + "_BaseDeleteFeatureMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_monitor._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10847,13 +10621,6 @@ def test_delete_feature_monitor_rest_required_fields(request_type=feature_regist assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_monitor_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_feature_monitor_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10954,17 +10721,20 @@ def test_create_feature_monitor_job_rest_required_fields(request_type=feature_re # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_monitor_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeatureMonitorJob, + "_BaseCreateFeatureMonitorJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature_monitor_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_monitor_job_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureMonitorJobId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11015,13 +10785,6 @@ def test_create_feature_monitor_job_rest_required_fields(request_type=feature_re assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_monitor_job_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature_monitor_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureMonitorJobId", )) & set(("parent", "featureMonitorJob", ))) - - def test_create_feature_monitor_job_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11128,16 +10891,18 @@ def test_get_feature_monitor_job_rest_required_fields(request_type=feature_regis # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_monitor_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeatureMonitorJob, + "_BaseGetFeatureMonitorJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature_monitor_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11186,13 +10951,6 @@ def test_get_feature_monitor_job_rest_required_fields(request_type=feature_regis assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_monitor_job_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature_monitor_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_feature_monitor_job_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11295,17 +11053,20 @@ def test_list_feature_monitor_jobs_rest_required_fields(request_type=feature_reg # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_monitor_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatureMonitorJobs, + "_BaseListFeatureMonitorJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_feature_monitor_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11355,13 +11116,6 @@ def test_list_feature_monitor_jobs_rest_required_fields(request_type=feature_reg assert sorted(expected_params) == sorted(actual_params) -def test_list_feature_monitor_jobs_rest_unset_required_fields(): - transport = transports.FeatureRegistryServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_feature_monitor_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_feature_monitor_jobs_rest_flattened(): client = FeatureRegistryServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_online_serving_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_online_serving_service.py index 770236482d..d66b0b5fa4 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_online_serving_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_online_serving_service.py @@ -81,6 +81,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -124,150 +136,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(None) is None - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeaturestoreOnlineServingServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeaturestoreOnlineServingServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreOnlineServingServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeaturestoreOnlineServingServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeaturestoreOnlineServingServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeaturestoreOnlineServingServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -281,40 +149,6 @@ def test__get_client_cert_source(): assert FeaturestoreOnlineServingServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeaturestoreOnlineServingServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeaturestoreOnlineServingServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreOnlineServingServiceClient)) -@mock.patch.object(FeaturestoreOnlineServingServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreOnlineServingServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeaturestoreOnlineServingServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeaturestoreOnlineServingServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreOnlineServingServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeaturestoreOnlineServingServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeaturestoreOnlineServingServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeaturestoreOnlineServingServiceClient._get_universe_domain(None, None) == FeaturestoreOnlineServingServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeaturestoreOnlineServingServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -699,11 +533,12 @@ def test_featurestore_online_serving_service_client_get_mtls_endpoint_and_cert_s for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -746,11 +581,12 @@ def test_featurestore_online_serving_service_client_get_mtls_endpoint_and_cert_s for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1961,7 +1797,12 @@ def test_read_feature_values_rest_required_fields(request_type=featurestore_onli # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadFeatureValues, + "_BaseReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -1969,9 +1810,6 @@ def test_read_feature_values_rest_required_fields(request_type=featurestore_onli jsonified_request["entityType"] = 'entity_type_value' jsonified_request["entityId"] = 'entity_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -2023,13 +1861,6 @@ def test_read_feature_values_rest_required_fields(request_type=featurestore_onli assert sorted(expected_params) == sorted(actual_params) -def test_read_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreOnlineServingServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "entityId", "featureSelector", ))) - - def test_read_feature_values_rest_flattened(): client = FeaturestoreOnlineServingServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2133,7 +1964,12 @@ def test_streaming_read_feature_values_rest_required_fields(request_type=feature # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).streaming_read_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamingReadFeatureValues, + "_BaseStreamingReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -2141,9 +1977,6 @@ def test_streaming_read_feature_values_rest_required_fields(request_type=feature jsonified_request["entityType"] = 'entity_type_value' jsonified_request["entityIds"] = 'entity_ids_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).streaming_read_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -2198,13 +2031,6 @@ def test_streaming_read_feature_values_rest_required_fields(request_type=feature assert sorted(expected_params) == sorted(actual_params) -def test_streaming_read_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreOnlineServingServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.streaming_read_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "entityIds", "featureSelector", ))) - - def test_streaming_read_feature_values_rest_flattened(): client = FeaturestoreOnlineServingServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2310,16 +2136,18 @@ def test_write_feature_values_rest_required_fields(request_type=featurestore_onl # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseWriteFeatureValues, + "_BaseWriteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -2369,13 +2197,6 @@ def test_write_feature_values_rest_required_fields(request_type=featurestore_onl assert sorted(expected_params) == sorted(actual_params) -def test_write_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreOnlineServingServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.write_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "payloads", ))) - - def test_write_feature_values_rest_flattened(): client = FeaturestoreOnlineServingServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_service.py index 24e674deb9..c054a324fb 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_featurestore_service.py @@ -102,6 +102,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -145,150 +157,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert FeaturestoreServiceClient._get_default_mtls_endpoint(None) is None - assert FeaturestoreServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert FeaturestoreServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert FeaturestoreServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - FeaturestoreServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert FeaturestoreServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert FeaturestoreServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - FeaturestoreServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert FeaturestoreServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -302,40 +170,6 @@ def test__get_client_cert_source(): assert FeaturestoreServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert FeaturestoreServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(FeaturestoreServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreServiceClient)) -@mock.patch.object(FeaturestoreServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(FeaturestoreServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = FeaturestoreServiceClient._DEFAULT_UNIVERSE - default_endpoint = FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = FeaturestoreServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert FeaturestoreServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert FeaturestoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert FeaturestoreServiceClient._get_api_endpoint(None, None, default_universe, "always") == FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == FeaturestoreServiceClient.DEFAULT_MTLS_ENDPOINT - assert FeaturestoreServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert FeaturestoreServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - FeaturestoreServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert FeaturestoreServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert FeaturestoreServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert FeaturestoreServiceClient._get_universe_domain(None, None) == FeaturestoreServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - FeaturestoreServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -720,11 +554,12 @@ def test_featurestore_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -767,11 +602,12 @@ def test_featurestore_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -8704,7 +8540,12 @@ def test_create_featurestore_rest_required_fields(request_type=featurestore_serv # verify fields with default values are dropped assert "featurestoreId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeaturestore, + "_BaseCreateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -8714,10 +8555,8 @@ def test_create_featurestore_rest_required_fields(request_type=featurestore_serv jsonified_request["parent"] = 'parent_value' jsonified_request["featurestoreId"] = 'featurestore_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_featurestore._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("featurestore_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featurestoreId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8771,13 +8610,6 @@ def test_create_featurestore_rest_required_fields(request_type=featurestore_serv assert sorted(expected_params) == sorted(actual_params) -def test_create_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featurestoreId", )) & set(("parent", "featurestore", "featurestoreId", ))) - - def test_create_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8882,16 +8714,18 @@ def test_get_featurestore_rest_required_fields(request_type=featurestore_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeaturestore, + "_BaseGetFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_featurestore._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8940,13 +8774,6 @@ def test_get_featurestore_rest_required_fields(request_type=featurestore_service assert sorted(expected_params) == sorted(actual_params) -def test_get_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9049,17 +8876,20 @@ def test_list_featurestores_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_featurestores._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeaturestores, + "_BaseListFeaturestores__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_featurestores._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9109,13 +8939,6 @@ def test_list_featurestores_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_list_featurestores_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_featurestores._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_featurestores_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9286,15 +9109,18 @@ def test_update_featurestore_rest_required_fields(request_type=featurestore_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeaturestore, + "_BaseUpdateFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_featurestore._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -9340,13 +9166,6 @@ def test_update_featurestore_rest_required_fields(request_type=featurestore_serv assert sorted(expected_params) == sorted(actual_params) -def test_update_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("featurestore", ))) - - def test_update_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9453,17 +9272,20 @@ def test_delete_featurestore_rest_required_fields(request_type=featurestore_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_featurestore._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeaturestore, + "_BaseDeleteFeaturestore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_featurestore._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9510,13 +9332,6 @@ def test_delete_featurestore_rest_required_fields(request_type=featurestore_serv assert sorted(expected_params) == sorted(actual_params) -def test_delete_featurestore_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_featurestore._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_featurestore_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9625,7 +9440,12 @@ def test_create_entity_type_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped assert "entityTypeId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateEntityType, + "_BaseCreateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -9635,10 +9455,8 @@ def test_create_entity_type_rest_required_fields(request_type=featurestore_servi jsonified_request["parent"] = 'parent_value' jsonified_request["entityTypeId"] = 'entity_type_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_entity_type._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("entity_type_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("entityTypeId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9692,13 +9510,6 @@ def test_create_entity_type_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_create_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(("entityTypeId", )) & set(("parent", "entityTypeId", ))) - - def test_create_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9803,16 +9614,18 @@ def test_get_entity_type_rest_required_fields(request_type=featurestore_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetEntityType, + "_BaseGetEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_entity_type._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9861,13 +9674,6 @@ def test_get_entity_type_rest_required_fields(request_type=featurestore_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9970,17 +9776,20 @@ def test_list_entity_types_rest_required_fields(request_type=featurestore_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_entity_types._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListEntityTypes, + "_BaseListEntityTypes__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_entity_types._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10030,13 +9839,6 @@ def test_list_entity_types_rest_required_fields(request_type=featurestore_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_entity_types_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_entity_types._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_entity_types_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10203,15 +10005,18 @@ def test_update_entity_type_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateEntityType, + "_BaseUpdateEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_entity_type._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -10260,13 +10065,6 @@ def test_update_entity_type_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_update_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("entityType", ))) - - def test_update_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10375,17 +10173,20 @@ def test_delete_entity_type_rest_required_fields(request_type=featurestore_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_entity_type._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteEntityType, + "_BaseDeleteEntityType__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_entity_type._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -10432,13 +10233,6 @@ def test_delete_entity_type_rest_required_fields(request_type=featurestore_servi assert sorted(expected_params) == sorted(actual_params) -def test_delete_entity_type_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_entity_type._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_entity_type_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10547,7 +10341,12 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C # verify fields with default values are dropped assert "featureId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateFeature, + "_BaseCreateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -10557,10 +10356,8 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C jsonified_request["parent"] = 'parent_value' jsonified_request["featureId"] = 'feature_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10614,13 +10411,6 @@ def test_create_feature_rest_required_fields(request_type=featurestore_service.C assert sorted(expected_params) == sorted(actual_params) -def test_create_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureId", )) & set(("parent", "feature", "featureId", ))) - - def test_create_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10729,16 +10519,18 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateFeatures, + "_BaseBatchCreateFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_features._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10785,13 +10577,6 @@ def test_batch_create_features_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_features_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_features_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10894,17 +10679,20 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetFeature, + "_BaseGetFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("feature_stats_and_anomaly_spec", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("featureStatsAndAnomalySpec", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -10954,13 +10742,6 @@ def test_get_feature_rest_required_fields(request_type=featurestore_service.GetF assert sorted(expected_params) == sorted(actual_params) -def test_get_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("featureStatsAndAnomalySpec", )) & set(("name", ))) - - def test_get_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11063,17 +10844,20 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListFeatures, + "_BaseListFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_features._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "latest_stats_count", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11123,13 +10907,6 @@ def test_list_features_rest_required_fields(request_type=featurestore_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_features_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "latestStatsCount", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_features_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11296,15 +11073,18 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateFeature, + "_BaseUpdateFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_feature._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -11353,13 +11133,6 @@ def test_update_feature_rest_required_fields(request_type=featurestore_service.U assert sorted(expected_params) == sorted(actual_params) -def test_update_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("feature", ))) - - def test_update_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11468,16 +11241,18 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeature, + "_BaseDeleteFeature__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11523,13 +11298,6 @@ def test_delete_feature_rest_required_fields(request_type=featurestore_service.D assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_feature_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11634,16 +11402,18 @@ def test_import_feature_values_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportFeatureValues, + "_BaseImportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -11690,13 +11460,6 @@ def test_import_feature_values_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_import_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "featureSpecs", ))) - - def test_import_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11801,16 +11564,18 @@ def test_batch_read_feature_values_rest_required_fields(request_type=featurestor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchReadFeatureValues, + "_BaseBatchReadFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["featurestore"] = 'featurestore_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "featurestore" in jsonified_request assert jsonified_request["featurestore"] == 'featurestore_value' @@ -11857,13 +11622,6 @@ def test_batch_read_feature_values_rest_required_fields(request_type=featurestor assert sorted(expected_params) == sorted(actual_params) -def test_batch_read_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_read_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("featurestore", "destination", "entityTypeSpecs", ))) - - def test_batch_read_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11968,16 +11726,18 @@ def test_export_feature_values_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportFeatureValues, + "_BaseExportFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -12024,13 +11784,6 @@ def test_export_feature_values_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_export_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", "destination", "featureSelector", ))) - - def test_export_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12135,16 +11888,18 @@ def test_delete_feature_values_rest_required_fields(request_type=featurestore_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_values._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteFeatureValues, + "_BaseDeleteFeatureValues__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["entityType"] = 'entity_type_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_feature_values._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "entityType" in jsonified_request assert jsonified_request["entityType"] == 'entity_type_value' @@ -12191,13 +11946,6 @@ def test_delete_feature_values_rest_required_fields(request_type=featurestore_se assert sorted(expected_params) == sorted(actual_params) -def test_delete_feature_values_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_feature_values._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("entityType", ))) - - def test_delete_feature_values_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12298,17 +12046,20 @@ def test_search_features_rest_required_fields(request_type=featurestore_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_features._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchFeatures, + "_BaseSearchFeatures__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["location"] = 'location_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_features._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", "query", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", "query", )) # verify required fields with non-default values are left alone assert "location" in jsonified_request @@ -12358,13 +12109,6 @@ def test_search_features_rest_required_fields(request_type=featurestore_service. assert sorted(expected_params) == sorted(actual_params) -def test_search_features_rest_unset_required_fields(): - transport = transports.FeaturestoreServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_features._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", "query", )) & set(("location", ))) - - def test_search_features_rest_flattened(): client = FeaturestoreServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_cache_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_cache_service.py index caa6a44d4d..0a42cd57b0 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_cache_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_cache_service.py @@ -90,6 +90,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -133,150 +145,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert GenAiCacheServiceClient._get_default_mtls_endpoint(None) is None - assert GenAiCacheServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert GenAiCacheServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert GenAiCacheServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiCacheServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - GenAiCacheServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert GenAiCacheServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiCacheServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert GenAiCacheServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - GenAiCacheServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert GenAiCacheServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -290,40 +158,6 @@ def test__get_client_cert_source(): assert GenAiCacheServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert GenAiCacheServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(GenAiCacheServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiCacheServiceClient)) -@mock.patch.object(GenAiCacheServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiCacheServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = GenAiCacheServiceClient._DEFAULT_UNIVERSE - default_endpoint = GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = GenAiCacheServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert GenAiCacheServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert GenAiCacheServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiCacheServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert GenAiCacheServiceClient._get_api_endpoint(None, None, default_universe, "always") == GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiCacheServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == GenAiCacheServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiCacheServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert GenAiCacheServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiCacheServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert GenAiCacheServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert GenAiCacheServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert GenAiCacheServiceClient._get_universe_domain(None, None) == GenAiCacheServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - GenAiCacheServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -708,11 +542,12 @@ def test_gen_ai_cache_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -755,11 +590,12 @@ def test_gen_ai_cache_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2823,16 +2659,18 @@ def test_create_cached_content_rest_required_fields(request_type=gen_ai_cache_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateCachedContent, + "_BaseCreateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_cached_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2882,13 +2720,6 @@ def test_create_cached_content_rest_required_fields(request_type=gen_ai_cache_se assert sorted(expected_params) == sorted(actual_params) -def test_create_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "cachedContent", ))) - - def test_create_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2993,16 +2824,18 @@ def test_get_cached_content_rest_required_fields(request_type=gen_ai_cache_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetCachedContent, + "_BaseGetCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_cached_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3051,13 +2884,6 @@ def test_get_cached_content_rest_required_fields(request_type=gen_ai_cache_servi assert sorted(expected_params) == sorted(actual_params) -def test_get_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3159,15 +2985,18 @@ def test_update_cached_content_rest_required_fields(request_type=gen_ai_cache_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateCachedContent, + "_BaseUpdateCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_cached_content._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3216,13 +3045,6 @@ def test_update_cached_content_rest_required_fields(request_type=gen_ai_cache_se assert sorted(expected_params) == sorted(actual_params) -def test_update_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("cachedContent", "updateMask", ))) - - def test_update_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3327,16 +3149,18 @@ def test_delete_cached_content_rest_required_fields(request_type=gen_ai_cache_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_cached_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteCachedContent, + "_BaseDeleteCachedContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_cached_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3382,13 +3206,6 @@ def test_delete_cached_content_rest_required_fields(request_type=gen_ai_cache_se assert sorted(expected_params) == sorted(actual_params) -def test_delete_cached_content_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_cached_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_cached_content_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3489,17 +3306,20 @@ def test_list_cached_contents_rest_required_fields(request_type=gen_ai_cache_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_cached_contents._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListCachedContents, + "_BaseListCachedContents__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_cached_contents._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3549,13 +3369,6 @@ def test_list_cached_contents_rest_required_fields(request_type=gen_ai_cache_ser assert sorted(expected_params) == sorted(actual_params) -def test_list_cached_contents_rest_unset_required_fields(): - transport = transports.GenAiCacheServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_cached_contents._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_cached_contents_rest_flattened(): client = GenAiCacheServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_tuning_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_tuning_service.py index 38ebe8358b..4aab9d8bdf 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_tuning_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_gen_ai_tuning_service.py @@ -97,6 +97,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -140,150 +152,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert GenAiTuningServiceClient._get_default_mtls_endpoint(None) is None - assert GenAiTuningServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert GenAiTuningServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert GenAiTuningServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiTuningServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - GenAiTuningServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert GenAiTuningServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiTuningServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert GenAiTuningServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - GenAiTuningServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert GenAiTuningServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -297,40 +165,6 @@ def test__get_client_cert_source(): assert GenAiTuningServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert GenAiTuningServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(GenAiTuningServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiTuningServiceClient)) -@mock.patch.object(GenAiTuningServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(GenAiTuningServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = GenAiTuningServiceClient._DEFAULT_UNIVERSE - default_endpoint = GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = GenAiTuningServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert GenAiTuningServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert GenAiTuningServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiTuningServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert GenAiTuningServiceClient._get_api_endpoint(None, None, default_universe, "always") == GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiTuningServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == GenAiTuningServiceClient.DEFAULT_MTLS_ENDPOINT - assert GenAiTuningServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert GenAiTuningServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - GenAiTuningServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert GenAiTuningServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert GenAiTuningServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert GenAiTuningServiceClient._get_universe_domain(None, None) == GenAiTuningServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - GenAiTuningServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -715,11 +549,12 @@ def test_gen_ai_tuning_service_client_get_mtls_endpoint_and_cert_source(client_c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -762,11 +597,12 @@ def test_gen_ai_tuning_service_client_get_mtls_endpoint_and_cert_source(client_c for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2883,16 +2719,18 @@ def test_create_tuning_job_rest_required_fields(request_type=genai_tuning_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTuningJob, + "_BaseCreateTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2942,13 +2780,6 @@ def test_create_tuning_job_rest_required_fields(request_type=genai_tuning_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_tuning_job_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "tuningJob", ))) - - def test_create_tuning_job_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3053,16 +2884,18 @@ def test_get_tuning_job_rest_required_fields(request_type=genai_tuning_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTuningJob, + "_BaseGetTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3111,13 +2944,6 @@ def test_get_tuning_job_rest_required_fields(request_type=genai_tuning_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_tuning_job_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tuning_job_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3220,17 +3046,20 @@ def test_list_tuning_jobs_rest_required_fields(request_type=genai_tuning_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tuning_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTuningJobs, + "_BaseListTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tuning_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3280,13 +3109,6 @@ def test_list_tuning_jobs_rest_required_fields(request_type=genai_tuning_service assert sorted(expected_params) == sorted(actual_params) -def test_list_tuning_jobs_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tuning_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_tuning_jobs_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3454,16 +3276,18 @@ def test_cancel_tuning_job_rest_required_fields(request_type=genai_tuning_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelTuningJob, + "_BaseCancelTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3510,13 +3334,6 @@ def test_cancel_tuning_job_rest_required_fields(request_type=genai_tuning_servic assert sorted(expected_params) == sorted(actual_params) -def test_cancel_tuning_job_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_tuning_job_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3621,16 +3438,18 @@ def test_rebase_tuned_model_rest_required_fields(request_type=genai_tuning_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).rebase_tuned_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRebaseTunedModel, + "_BaseRebaseTunedModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).rebase_tuned_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3677,13 +3496,6 @@ def test_rebase_tuned_model_rest_required_fields(request_type=genai_tuning_servi assert sorted(expected_params) == sorted(actual_params) -def test_rebase_tuned_model_rest_unset_required_fields(): - transport = transports.GenAiTuningServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.rebase_tuned_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "tunedModelRef", ))) - - def test_rebase_tuned_model_rest_flattened(): client = GenAiTuningServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_index_endpoint_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_index_endpoint_service.py index 705eec115c..e1ff717819 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_index_endpoint_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_index_endpoint_service.py @@ -95,6 +95,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -138,150 +150,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert IndexEndpointServiceClient._get_default_mtls_endpoint(None) is None - assert IndexEndpointServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert IndexEndpointServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert IndexEndpointServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexEndpointServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - IndexEndpointServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert IndexEndpointServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexEndpointServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert IndexEndpointServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - IndexEndpointServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert IndexEndpointServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -295,40 +163,6 @@ def test__get_client_cert_source(): assert IndexEndpointServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert IndexEndpointServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(IndexEndpointServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexEndpointServiceClient)) -@mock.patch.object(IndexEndpointServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexEndpointServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = IndexEndpointServiceClient._DEFAULT_UNIVERSE - default_endpoint = IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = IndexEndpointServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert IndexEndpointServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert IndexEndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexEndpointServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert IndexEndpointServiceClient._get_api_endpoint(None, None, default_universe, "always") == IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexEndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == IndexEndpointServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexEndpointServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert IndexEndpointServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexEndpointServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert IndexEndpointServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert IndexEndpointServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert IndexEndpointServiceClient._get_universe_domain(None, None) == IndexEndpointServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - IndexEndpointServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -713,11 +547,12 @@ def test_index_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_ for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -760,11 +595,12 @@ def test_index_endpoint_service_client_get_mtls_endpoint_and_cert_source(client_ for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3886,16 +3722,18 @@ def test_create_index_endpoint_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateIndexEndpoint, + "_BaseCreateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3942,13 +3780,6 @@ def test_create_index_endpoint_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_create_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "indexEndpoint", ))) - - def test_create_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4051,16 +3882,18 @@ def test_get_index_endpoint_rest_required_fields(request_type=index_endpoint_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetIndexEndpoint, + "_BaseGetIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4109,13 +3942,6 @@ def test_get_index_endpoint_rest_required_fields(request_type=index_endpoint_ser assert sorted(expected_params) == sorted(actual_params) -def test_get_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4218,17 +4044,20 @@ def test_list_index_endpoints_rest_required_fields(request_type=index_endpoint_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_index_endpoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListIndexEndpoints, + "_BaseListIndexEndpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_index_endpoints._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4278,13 +4107,6 @@ def test_list_index_endpoints_rest_required_fields(request_type=index_endpoint_s assert sorted(expected_params) == sorted(actual_params) -def test_list_index_endpoints_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_index_endpoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_index_endpoints_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4451,15 +4273,18 @@ def test_update_index_endpoint_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateIndexEndpoint, + "_BaseUpdateIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index_endpoint._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4508,13 +4333,6 @@ def test_update_index_endpoint_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_update_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("indexEndpoint", "updateMask", ))) - - def test_update_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4623,16 +4441,18 @@ def test_delete_index_endpoint_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index_endpoint._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteIndexEndpoint, + "_BaseDeleteIndexEndpoint__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index_endpoint._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4678,13 +4498,6 @@ def test_delete_index_endpoint_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_delete_index_endpoint_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_index_endpoint._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_index_endpoint_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4789,16 +4602,18 @@ def test_deploy_index_rest_required_fields(request_type=index_endpoint_service.D # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeployIndex, + "_BaseDeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -4845,13 +4660,6 @@ def test_deploy_index_rest_required_fields(request_type=index_endpoint_service.D assert sorted(expected_params) == sorted(actual_params) -def test_deploy_index_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.deploy_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", "deployedIndex", ))) - - def test_deploy_index_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4959,7 +4767,12 @@ def test_undeploy_index_rest_required_fields(request_type=index_endpoint_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUndeployIndex, + "_BaseUndeployIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -4967,9 +4780,6 @@ def test_undeploy_index_rest_required_fields(request_type=index_endpoint_service jsonified_request["indexEndpoint"] = 'index_endpoint_value' jsonified_request["deployedIndexId"] = 'deployed_index_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).undeploy_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -5018,13 +4828,6 @@ def test_undeploy_index_rest_required_fields(request_type=index_endpoint_service assert sorted(expected_params) == sorted(actual_params) -def test_undeploy_index_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.undeploy_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", "deployedIndexId", ))) - - def test_undeploy_index_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5131,16 +4934,18 @@ def test_mutate_deployed_index_rest_required_fields(request_type=index_endpoint_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseMutateDeployedIndex, + "_BaseMutateDeployedIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).mutate_deployed_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -5187,13 +4992,6 @@ def test_mutate_deployed_index_rest_required_fields(request_type=index_endpoint_ assert sorted(expected_params) == sorted(actual_params) -def test_mutate_deployed_index_rest_unset_required_fields(): - transport = transports.IndexEndpointServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.mutate_deployed_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", "deployedIndex", ))) - - def test_mutate_deployed_index_rest_flattened(): client = IndexEndpointServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_index_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_index_service.py index 5de7c36ea8..ded944c249 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_index_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_index_service.py @@ -92,6 +92,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -135,150 +147,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert IndexServiceClient._get_default_mtls_endpoint(None) is None - assert IndexServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert IndexServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert IndexServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert IndexServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - IndexServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert IndexServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert IndexServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert IndexServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert IndexServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert IndexServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert IndexServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - IndexServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert IndexServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert IndexServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -292,40 +160,6 @@ def test__get_client_cert_source(): assert IndexServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert IndexServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(IndexServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexServiceClient)) -@mock.patch.object(IndexServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(IndexServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = IndexServiceClient._DEFAULT_UNIVERSE - default_endpoint = IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = IndexServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert IndexServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert IndexServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == IndexServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert IndexServiceClient._get_api_endpoint(None, None, default_universe, "always") == IndexServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == IndexServiceClient.DEFAULT_MTLS_ENDPOINT - assert IndexServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert IndexServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - IndexServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert IndexServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert IndexServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert IndexServiceClient._get_universe_domain(None, None) == IndexServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - IndexServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -710,11 +544,12 @@ def test_index_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -757,11 +592,12 @@ def test_index_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3543,16 +3379,18 @@ def test_create_index_rest_required_fields(request_type=index_service.CreateInde # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateIndex, + "_BaseCreateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3599,13 +3437,6 @@ def test_create_index_rest_required_fields(request_type=index_service.CreateInde assert sorted(expected_params) == sorted(actual_params) -def test_create_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "index", ))) - - def test_create_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3708,16 +3539,18 @@ def test_get_index_rest_required_fields(request_type=index_service.GetIndexReque # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetIndex, + "_BaseGetIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3766,13 +3599,6 @@ def test_get_index_rest_required_fields(request_type=index_service.GetIndexReque assert sorted(expected_params) == sorted(actual_params) -def test_get_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3879,16 +3705,18 @@ def test_import_index_rest_required_fields(request_type=index_service.ImportInde # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportIndex, + "_BaseImportIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3935,13 +3763,6 @@ def test_import_index_rest_required_fields(request_type=index_service.ImportInde assert sorted(expected_params) == sorted(actual_params) -def test_import_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "config", ))) - - def test_list_indexes_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -3990,17 +3811,20 @@ def test_list_indexes_rest_required_fields(request_type=index_service.ListIndexe # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_indexes._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListIndexes, + "_BaseListIndexes__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_indexes._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4050,13 +3874,6 @@ def test_list_indexes_rest_required_fields(request_type=index_service.ListIndexe assert sorted(expected_params) == sorted(actual_params) -def test_list_indexes_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_indexes._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_indexes_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4227,15 +4044,18 @@ def test_update_index_rest_required_fields(request_type=index_service.UpdateInde # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateIndex, + "_BaseUpdateIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_index._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4281,13 +4101,6 @@ def test_update_index_rest_required_fields(request_type=index_service.UpdateInde assert sorted(expected_params) == sorted(actual_params) -def test_update_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("index", ))) - - def test_update_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4394,16 +4207,18 @@ def test_delete_index_rest_required_fields(request_type=index_service.DeleteInde # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteIndex, + "_BaseDeleteIndex__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_index._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4449,13 +4264,6 @@ def test_delete_index_rest_required_fields(request_type=index_service.DeleteInde assert sorted(expected_params) == sorted(actual_params) -def test_delete_index_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_index._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_index_rest_flattened(): client = IndexServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4556,16 +4364,18 @@ def test_upsert_datapoints_rest_required_fields(request_type=index_service.Upser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upsert_datapoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpsertDatapoints, + "_BaseUpsertDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["index"] = 'index_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upsert_datapoints._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "index" in jsonified_request assert jsonified_request["index"] == 'index_value' @@ -4615,13 +4425,6 @@ def test_upsert_datapoints_rest_required_fields(request_type=index_service.Upser assert sorted(expected_params) == sorted(actual_params) -def test_upsert_datapoints_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upsert_datapoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("index", ))) - - def test_remove_datapoints_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -4670,16 +4473,18 @@ def test_remove_datapoints_rest_required_fields(request_type=index_service.Remov # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_datapoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRemoveDatapoints, + "_BaseRemoveDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["index"] = 'index_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_datapoints._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "index" in jsonified_request assert jsonified_request["index"] == 'index_value' @@ -4729,13 +4534,6 @@ def test_remove_datapoints_rest_required_fields(request_type=index_service.Remov assert sorted(expected_params) == sorted(actual_params) -def test_remove_datapoints_rest_unset_required_fields(): - transport = transports.IndexServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.remove_datapoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("index", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.IndexServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_job_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_job_service.py index b60da331dd..d45507eeb0 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_job_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_job_service.py @@ -122,6 +122,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -165,150 +177,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert JobServiceClient._get_default_mtls_endpoint(None) is None - assert JobServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert JobServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert JobServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert JobServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert JobServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert JobServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - JobServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert JobServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert JobServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert JobServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert JobServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - JobServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert JobServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert JobServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - JobServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert JobServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert JobServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -322,40 +190,6 @@ def test__get_client_cert_source(): assert JobServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert JobServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(JobServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(JobServiceClient)) -@mock.patch.object(JobServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(JobServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = JobServiceClient._DEFAULT_UNIVERSE - default_endpoint = JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = JobServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert JobServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert JobServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == JobServiceClient.DEFAULT_MTLS_ENDPOINT - assert JobServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert JobServiceClient._get_api_endpoint(None, None, default_universe, "always") == JobServiceClient.DEFAULT_MTLS_ENDPOINT - assert JobServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == JobServiceClient.DEFAULT_MTLS_ENDPOINT - assert JobServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert JobServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - JobServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert JobServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert JobServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert JobServiceClient._get_universe_domain(None, None) == JobServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - JobServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -740,11 +574,12 @@ def test_job_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -787,11 +622,12 @@ def test_job_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -13954,16 +13790,18 @@ def test_create_custom_job_rest_required_fields(request_type=job_service.CreateC # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateCustomJob, + "_BaseCreateCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14013,13 +13851,6 @@ def test_create_custom_job_rest_required_fields(request_type=job_service.CreateC assert sorted(expected_params) == sorted(actual_params) -def test_create_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "customJob", ))) - - def test_create_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14124,16 +13955,18 @@ def test_get_custom_job_rest_required_fields(request_type=job_service.GetCustomJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetCustomJob, + "_BaseGetCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14182,13 +14015,6 @@ def test_get_custom_job_rest_required_fields(request_type=job_service.GetCustomJ assert sorted(expected_params) == sorted(actual_params) -def test_get_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14291,17 +14117,20 @@ def test_list_custom_jobs_rest_required_fields(request_type=job_service.ListCust # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_custom_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListCustomJobs, + "_BaseListCustomJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_custom_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14351,13 +14180,6 @@ def test_list_custom_jobs_rest_required_fields(request_type=job_service.ListCust assert sorted(expected_params) == sorted(actual_params) -def test_list_custom_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_custom_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_custom_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14529,16 +14351,18 @@ def test_delete_custom_job_rest_required_fields(request_type=job_service.DeleteC # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteCustomJob, + "_BaseDeleteCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14584,13 +14408,6 @@ def test_delete_custom_job_rest_required_fields(request_type=job_service.DeleteC assert sorted(expected_params) == sorted(actual_params) -def test_delete_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14691,16 +14508,18 @@ def test_cancel_custom_job_rest_required_fields(request_type=job_service.CancelC # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_custom_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelCustomJob, + "_BaseCancelCustomJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_custom_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14747,13 +14566,6 @@ def test_cancel_custom_job_rest_required_fields(request_type=job_service.CancelC assert sorted(expected_params) == sorted(actual_params) -def test_cancel_custom_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_custom_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_custom_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14854,16 +14666,18 @@ def test_create_data_labeling_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateDataLabelingJob, + "_BaseCreateDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14913,13 +14727,6 @@ def test_create_data_labeling_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_create_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "dataLabelingJob", ))) - - def test_create_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15024,16 +14831,18 @@ def test_get_data_labeling_job_rest_required_fields(request_type=job_service.Get # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetDataLabelingJob, + "_BaseGetDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15082,13 +14891,6 @@ def test_get_data_labeling_job_rest_required_fields(request_type=job_service.Get assert sorted(expected_params) == sorted(actual_params) -def test_get_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15191,17 +14993,20 @@ def test_list_data_labeling_jobs_rest_required_fields(request_type=job_service.L # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_labeling_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListDataLabelingJobs, + "_BaseListDataLabelingJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_data_labeling_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -15251,13 +15056,6 @@ def test_list_data_labeling_jobs_rest_required_fields(request_type=job_service.L assert sorted(expected_params) == sorted(actual_params) -def test_list_data_labeling_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_data_labeling_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_data_labeling_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15429,16 +15227,18 @@ def test_delete_data_labeling_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteDataLabelingJob, + "_BaseDeleteDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15484,13 +15284,6 @@ def test_delete_data_labeling_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_delete_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15591,16 +15384,18 @@ def test_cancel_data_labeling_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_data_labeling_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelDataLabelingJob, + "_BaseCancelDataLabelingJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_data_labeling_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15647,13 +15442,6 @@ def test_cancel_data_labeling_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_cancel_data_labeling_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_data_labeling_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_data_labeling_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15754,16 +15542,18 @@ def test_create_hyperparameter_tuning_job_rest_required_fields(request_type=job_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateHyperparameterTuningJob, + "_BaseCreateHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -15813,13 +15603,6 @@ def test_create_hyperparameter_tuning_job_rest_required_fields(request_type=job_ assert sorted(expected_params) == sorted(actual_params) -def test_create_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "hyperparameterTuningJob", ))) - - def test_create_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15924,16 +15707,18 @@ def test_get_hyperparameter_tuning_job_rest_required_fields(request_type=job_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetHyperparameterTuningJob, + "_BaseGetHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15982,13 +15767,6 @@ def test_get_hyperparameter_tuning_job_rest_required_fields(request_type=job_ser assert sorted(expected_params) == sorted(actual_params) -def test_get_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16091,17 +15869,20 @@ def test_list_hyperparameter_tuning_jobs_rest_required_fields(request_type=job_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_hyperparameter_tuning_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListHyperparameterTuningJobs, + "_BaseListHyperparameterTuningJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_hyperparameter_tuning_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -16151,13 +15932,6 @@ def test_list_hyperparameter_tuning_jobs_rest_required_fields(request_type=job_s assert sorted(expected_params) == sorted(actual_params) -def test_list_hyperparameter_tuning_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_hyperparameter_tuning_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_hyperparameter_tuning_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16329,16 +16103,18 @@ def test_delete_hyperparameter_tuning_job_rest_required_fields(request_type=job_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteHyperparameterTuningJob, + "_BaseDeleteHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16384,13 +16160,6 @@ def test_delete_hyperparameter_tuning_job_rest_required_fields(request_type=job_ assert sorted(expected_params) == sorted(actual_params) -def test_delete_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16491,16 +16260,18 @@ def test_cancel_hyperparameter_tuning_job_rest_required_fields(request_type=job_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelHyperparameterTuningJob, + "_BaseCancelHyperparameterTuningJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_hyperparameter_tuning_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16547,13 +16318,6 @@ def test_cancel_hyperparameter_tuning_job_rest_required_fields(request_type=job_ assert sorted(expected_params) == sorted(actual_params) -def test_cancel_hyperparameter_tuning_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_hyperparameter_tuning_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_hyperparameter_tuning_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16654,16 +16418,18 @@ def test_create_nas_job_rest_required_fields(request_type=job_service.CreateNasJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateNasJob, + "_BaseCreateNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -16713,13 +16479,6 @@ def test_create_nas_job_rest_required_fields(request_type=job_service.CreateNasJ assert sorted(expected_params) == sorted(actual_params) -def test_create_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "nasJob", ))) - - def test_create_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16824,16 +16583,18 @@ def test_get_nas_job_rest_required_fields(request_type=job_service.GetNasJobRequ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNasJob, + "_BaseGetNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16882,13 +16643,6 @@ def test_get_nas_job_rest_required_fields(request_type=job_service.GetNasJobRequ assert sorted(expected_params) == sorted(actual_params) -def test_get_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16991,17 +16745,20 @@ def test_list_nas_jobs_rest_required_fields(request_type=job_service.ListNasJobs # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNasJobs, + "_BaseListNasJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17051,13 +16808,6 @@ def test_list_nas_jobs_rest_required_fields(request_type=job_service.ListNasJobs assert sorted(expected_params) == sorted(actual_params) -def test_list_nas_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_nas_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_nas_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17229,16 +16979,18 @@ def test_delete_nas_job_rest_required_fields(request_type=job_service.DeleteNasJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNasJob, + "_BaseDeleteNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17284,13 +17036,6 @@ def test_delete_nas_job_rest_required_fields(request_type=job_service.DeleteNasJ assert sorted(expected_params) == sorted(actual_params) -def test_delete_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17391,16 +17136,18 @@ def test_cancel_nas_job_rest_required_fields(request_type=job_service.CancelNasJ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_nas_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelNasJob, + "_BaseCancelNasJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_nas_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17447,13 +17194,6 @@ def test_cancel_nas_job_rest_required_fields(request_type=job_service.CancelNasJ assert sorted(expected_params) == sorted(actual_params) -def test_cancel_nas_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_nas_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_nas_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17554,16 +17294,18 @@ def test_get_nas_trial_detail_rest_required_fields(request_type=job_service.GetN # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_trial_detail._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNasTrialDetail, + "_BaseGetNasTrialDetail__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_nas_trial_detail._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17612,13 +17354,6 @@ def test_get_nas_trial_detail_rest_required_fields(request_type=job_service.GetN assert sorted(expected_params) == sorted(actual_params) -def test_get_nas_trial_detail_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_nas_trial_detail._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_nas_trial_detail_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17721,17 +17456,20 @@ def test_list_nas_trial_details_rest_required_fields(request_type=job_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_trial_details._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNasTrialDetails, + "_BaseListNasTrialDetails__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_nas_trial_details._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17781,13 +17519,6 @@ def test_list_nas_trial_details_rest_required_fields(request_type=job_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_nas_trial_details_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_nas_trial_details._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_nas_trial_details_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17955,16 +17686,18 @@ def test_create_batch_prediction_job_rest_required_fields(request_type=job_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateBatchPredictionJob, + "_BaseCreateBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -18014,13 +17747,6 @@ def test_create_batch_prediction_job_rest_required_fields(request_type=job_servi assert sorted(expected_params) == sorted(actual_params) -def test_create_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "batchPredictionJob", ))) - - def test_create_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18125,16 +17851,18 @@ def test_get_batch_prediction_job_rest_required_fields(request_type=job_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetBatchPredictionJob, + "_BaseGetBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -18183,13 +17911,6 @@ def test_get_batch_prediction_job_rest_required_fields(request_type=job_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18292,17 +18013,20 @@ def test_list_batch_prediction_jobs_rest_required_fields(request_type=job_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_batch_prediction_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListBatchPredictionJobs, + "_BaseListBatchPredictionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_batch_prediction_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -18352,13 +18076,6 @@ def test_list_batch_prediction_jobs_rest_required_fields(request_type=job_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_batch_prediction_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_batch_prediction_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_batch_prediction_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18530,16 +18247,18 @@ def test_delete_batch_prediction_job_rest_required_fields(request_type=job_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteBatchPredictionJob, + "_BaseDeleteBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -18585,13 +18304,6 @@ def test_delete_batch_prediction_job_rest_required_fields(request_type=job_servi assert sorted(expected_params) == sorted(actual_params) -def test_delete_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18692,16 +18404,18 @@ def test_cancel_batch_prediction_job_rest_required_fields(request_type=job_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_batch_prediction_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelBatchPredictionJob, + "_BaseCancelBatchPredictionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_batch_prediction_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -18748,13 +18462,6 @@ def test_cancel_batch_prediction_job_rest_required_fields(request_type=job_servi assert sorted(expected_params) == sorted(actual_params) -def test_cancel_batch_prediction_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_batch_prediction_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_batch_prediction_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18855,16 +18562,18 @@ def test_create_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateModelDeploymentMonitoringJob, + "_BaseCreateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -18914,13 +18623,6 @@ def test_create_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_create_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "modelDeploymentMonitoringJob", ))) - - def test_create_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19026,7 +18728,12 @@ def test_search_model_deployment_monitoring_stats_anomalies_rest_required_fields # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_deployment_monitoring_stats_anomalies._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchModelDeploymentMonitoringStatsAnomalies, + "_BaseSearchModelDeploymentMonitoringStatsAnomalies__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -19034,9 +18741,6 @@ def test_search_model_deployment_monitoring_stats_anomalies_rest_required_fields jsonified_request["modelDeploymentMonitoringJob"] = 'model_deployment_monitoring_job_value' jsonified_request["deployedModelId"] = 'deployed_model_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_deployment_monitoring_stats_anomalies._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "modelDeploymentMonitoringJob" in jsonified_request assert jsonified_request["modelDeploymentMonitoringJob"] == 'model_deployment_monitoring_job_value' @@ -19088,13 +18792,6 @@ def test_search_model_deployment_monitoring_stats_anomalies_rest_required_fields assert sorted(expected_params) == sorted(actual_params) -def test_search_model_deployment_monitoring_stats_anomalies_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_model_deployment_monitoring_stats_anomalies._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("modelDeploymentMonitoringJob", "deployedModelId", "objectives", ))) - - def test_search_model_deployment_monitoring_stats_anomalies_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19264,16 +18961,18 @@ def test_get_model_deployment_monitoring_job_rest_required_fields(request_type=j # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelDeploymentMonitoringJob, + "_BaseGetModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -19322,13 +19021,6 @@ def test_get_model_deployment_monitoring_job_rest_required_fields(request_type=j assert sorted(expected_params) == sorted(actual_params) -def test_get_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19431,17 +19123,20 @@ def test_list_model_deployment_monitoring_jobs_rest_required_fields(request_type # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_deployment_monitoring_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelDeploymentMonitoringJobs, + "_BaseListModelDeploymentMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_deployment_monitoring_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -19491,13 +19186,6 @@ def test_list_model_deployment_monitoring_jobs_rest_required_fields(request_type assert sorted(expected_params) == sorted(actual_params) -def test_list_model_deployment_monitoring_jobs_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_deployment_monitoring_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_deployment_monitoring_jobs_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19668,15 +19356,18 @@ def test_update_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateModelDeploymentMonitoringJob, + "_BaseUpdateModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -19722,13 +19413,6 @@ def test_update_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_update_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("modelDeploymentMonitoringJob", "updateMask", ))) - - def test_update_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19835,16 +19519,18 @@ def test_delete_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModelDeploymentMonitoringJob, + "_BaseDeleteModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -19890,13 +19576,6 @@ def test_delete_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -19997,16 +19676,18 @@ def test_pause_model_deployment_monitoring_job_rest_required_fields(request_type # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePauseModelDeploymentMonitoringJob, + "_BasePauseModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -20053,13 +19734,6 @@ def test_pause_model_deployment_monitoring_job_rest_required_fields(request_type assert sorted(expected_params) == sorted(actual_params) -def test_pause_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.pause_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_pause_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -20160,16 +19834,18 @@ def test_resume_model_deployment_monitoring_job_rest_required_fields(request_typ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseResumeModelDeploymentMonitoringJob, + "_BaseResumeModelDeploymentMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_model_deployment_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -20216,13 +19892,6 @@ def test_resume_model_deployment_monitoring_job_rest_required_fields(request_typ assert sorted(expected_params) == sorted(actual_params) -def test_resume_model_deployment_monitoring_job_rest_unset_required_fields(): - transport = transports.JobServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.resume_model_deployment_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_resume_model_deployment_monitoring_job_rest_flattened(): client = JobServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_llm_utility_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_llm_utility_service.py index 1b0c6151e0..a71a256819 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_llm_utility_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_llm_utility_service.py @@ -82,6 +82,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -125,150 +137,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert LlmUtilityServiceClient._get_default_mtls_endpoint(None) is None - assert LlmUtilityServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert LlmUtilityServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert LlmUtilityServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert LlmUtilityServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - LlmUtilityServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert LlmUtilityServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - LlmUtilityServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert LlmUtilityServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - LlmUtilityServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert LlmUtilityServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -282,40 +150,6 @@ def test__get_client_cert_source(): assert LlmUtilityServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert LlmUtilityServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(LlmUtilityServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(LlmUtilityServiceClient)) -@mock.patch.object(LlmUtilityServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(LlmUtilityServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = LlmUtilityServiceClient._DEFAULT_UNIVERSE - default_endpoint = LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = LlmUtilityServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert LlmUtilityServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert LlmUtilityServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - assert LlmUtilityServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert LlmUtilityServiceClient._get_api_endpoint(None, None, default_universe, "always") == LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - assert LlmUtilityServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == LlmUtilityServiceClient.DEFAULT_MTLS_ENDPOINT - assert LlmUtilityServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert LlmUtilityServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - LlmUtilityServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert LlmUtilityServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert LlmUtilityServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert LlmUtilityServiceClient._get_universe_domain(None, None) == LlmUtilityServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - LlmUtilityServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -700,11 +534,12 @@ def test_llm_utility_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -747,11 +582,12 @@ def test_llm_utility_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1349,16 +1185,18 @@ def test_compute_tokens_rest_required_fields(request_type=llm_utility_service.Co # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).compute_tokens._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseComputeTokens, + "_BaseComputeTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).compute_tokens._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -1408,13 +1246,6 @@ def test_compute_tokens_rest_required_fields(request_type=llm_utility_service.Co assert sorted(expected_params) == sorted(actual_params) -def test_compute_tokens_rest_unset_required_fields(): - transport = transports.LlmUtilityServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.compute_tokens._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_compute_tokens_rest_flattened(): client = LlmUtilityServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_match_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_match_service.py index 77f2ae4bf1..49b720808c 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_match_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_match_service.py @@ -80,6 +80,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -123,150 +135,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert MatchServiceClient._get_default_mtls_endpoint(None) is None - assert MatchServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert MatchServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert MatchServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert MatchServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MatchServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MatchServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - MatchServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert MatchServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert MatchServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert MatchServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert MatchServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - MatchServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert MatchServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert MatchServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - MatchServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert MatchServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert MatchServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -280,40 +148,6 @@ def test__get_client_cert_source(): assert MatchServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert MatchServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(MatchServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MatchServiceClient)) -@mock.patch.object(MatchServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MatchServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = MatchServiceClient._DEFAULT_UNIVERSE - default_endpoint = MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = MatchServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert MatchServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert MatchServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == MatchServiceClient.DEFAULT_MTLS_ENDPOINT - assert MatchServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert MatchServiceClient._get_api_endpoint(None, None, default_universe, "always") == MatchServiceClient.DEFAULT_MTLS_ENDPOINT - assert MatchServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == MatchServiceClient.DEFAULT_MTLS_ENDPOINT - assert MatchServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert MatchServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - MatchServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert MatchServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert MatchServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert MatchServiceClient._get_universe_domain(None, None) == MatchServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - MatchServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -698,11 +532,12 @@ def test_match_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -745,11 +580,12 @@ def test_match_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1480,16 +1316,18 @@ def test_find_neighbors_rest_required_fields(request_type=match_service.FindNeig # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).find_neighbors._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseFindNeighbors, + "_BaseFindNeighbors__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).find_neighbors._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -1539,13 +1377,6 @@ def test_find_neighbors_rest_required_fields(request_type=match_service.FindNeig assert sorted(expected_params) == sorted(actual_params) -def test_find_neighbors_rest_unset_required_fields(): - transport = transports.MatchServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.find_neighbors._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", ))) - - def test_read_index_datapoints_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -1594,16 +1425,18 @@ def test_read_index_datapoints_rest_required_fields(request_type=match_service.R # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_index_datapoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadIndexDatapoints, + "_BaseReadIndexDatapoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["indexEndpoint"] = 'index_endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_index_datapoints._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "indexEndpoint" in jsonified_request assert jsonified_request["indexEndpoint"] == 'index_endpoint_value' @@ -1653,13 +1486,6 @@ def test_read_index_datapoints_rest_required_fields(request_type=match_service.R assert sorted(expected_params) == sorted(actual_params) -def test_read_index_datapoints_rest_unset_required_fields(): - transport = transports.MatchServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_index_datapoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("indexEndpoint", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.MatchServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_memory_bank_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_memory_bank_service.py index 42c2a64847..8257e5a3ad 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_memory_bank_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_memory_bank_service.py @@ -91,6 +91,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -134,150 +146,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert MemoryBankServiceClient._get_default_mtls_endpoint(None) is None - assert MemoryBankServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert MemoryBankServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert MemoryBankServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert MemoryBankServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert MemoryBankServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert MemoryBankServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert MemoryBankServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MemoryBankServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MemoryBankServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - MemoryBankServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert MemoryBankServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert MemoryBankServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert MemoryBankServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert MemoryBankServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - MemoryBankServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert MemoryBankServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert MemoryBankServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert MemoryBankServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MemoryBankServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MemoryBankServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert MemoryBankServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert MemoryBankServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert MemoryBankServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert MemoryBankServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert MemoryBankServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - MemoryBankServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert MemoryBankServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert MemoryBankServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -291,40 +159,6 @@ def test__get_client_cert_source(): assert MemoryBankServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert MemoryBankServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(MemoryBankServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MemoryBankServiceClient)) -@mock.patch.object(MemoryBankServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MemoryBankServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = MemoryBankServiceClient._DEFAULT_UNIVERSE - default_endpoint = MemoryBankServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = MemoryBankServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert MemoryBankServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert MemoryBankServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == MemoryBankServiceClient.DEFAULT_MTLS_ENDPOINT - assert MemoryBankServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert MemoryBankServiceClient._get_api_endpoint(None, None, default_universe, "always") == MemoryBankServiceClient.DEFAULT_MTLS_ENDPOINT - assert MemoryBankServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == MemoryBankServiceClient.DEFAULT_MTLS_ENDPOINT - assert MemoryBankServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert MemoryBankServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - MemoryBankServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert MemoryBankServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert MemoryBankServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert MemoryBankServiceClient._get_universe_domain(None, None) == MemoryBankServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - MemoryBankServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -709,11 +543,12 @@ def test_memory_bank_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -756,11 +591,12 @@ def test_memory_bank_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1356,6 +1192,9 @@ def test_get_memory(request_type, transport: str = 'grpc'): display_name='display_name_value', description='description_value', fact='fact_value', + memory_type=memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION, + context='context_value', + disable_memory_revisions=True, ) response = client.get_memory(request) @@ -1371,6 +1210,8 @@ def test_get_memory(request_type, transport: str = 'grpc'): assert response.display_name == 'display_name_value' assert response.description == 'description_value' assert response.fact == 'fact_value' + assert response.memory_type == memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION + assert response.context == 'context_value' def test_get_memory_non_empty_request_with_auto_populated_field(): @@ -1492,6 +1333,8 @@ async def test_get_memory_async(request_type, transport: str = 'grpc_asyncio'): display_name='display_name_value', description='description_value', fact='fact_value', + memory_type=memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION, + context='context_value', )) response = await client.get_memory(request) @@ -1507,6 +1350,8 @@ async def test_get_memory_async(request_type, transport: str = 'grpc_asyncio'): assert response.display_name == 'display_name_value' assert response.description == 'description_value' assert response.fact == 'fact_value' + assert response.memory_type == memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION + assert response.context == 'context_value' def test_get_memory_field_headers(): client = MemoryBankServiceClient( @@ -3486,17 +3331,20 @@ def test_create_memory_rest_required_fields(request_type=memory_bank_service.Cre # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_memory._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateMemory, + "_BaseCreateMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_memory._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("memory_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("memoryId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3544,13 +3392,6 @@ def test_create_memory_rest_required_fields(request_type=memory_bank_service.Cre assert sorted(expected_params) == sorted(actual_params) -def test_create_memory_rest_unset_required_fields(): - transport = transports.MemoryBankServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_memory._get_unset_required_fields({}) - assert set(unset_fields) == (set(("memoryId", )) & set(("parent", "memory", ))) - - def test_create_memory_rest_flattened(): client = MemoryBankServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3655,16 +3496,18 @@ def test_get_memory_rest_required_fields(request_type=memory_bank_service.GetMem # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_memory._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetMemory, + "_BaseGetMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_memory._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3713,13 +3556,6 @@ def test_get_memory_rest_required_fields(request_type=memory_bank_service.GetMem assert sorted(expected_params) == sorted(actual_params) -def test_get_memory_rest_unset_required_fields(): - transport = transports.MemoryBankServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_memory._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_memory_rest_flattened(): client = MemoryBankServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3825,15 +3661,18 @@ def test_update_memory_rest_required_fields(request_type=memory_bank_service.Upd # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_memory._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateMemory, + "_BaseUpdateMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_memory._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3879,13 +3718,6 @@ def test_update_memory_rest_required_fields(request_type=memory_bank_service.Upd assert sorted(expected_params) == sorted(actual_params) -def test_update_memory_rest_unset_required_fields(): - transport = transports.MemoryBankServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_memory._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("memory", ))) - - def test_update_memory_rest_flattened(): client = MemoryBankServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3988,17 +3820,20 @@ def test_list_memories_rest_required_fields(request_type=memory_bank_service.Lis # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_memories._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListMemories, + "_BaseListMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_memories._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4048,13 +3883,6 @@ def test_list_memories_rest_required_fields(request_type=memory_bank_service.Lis assert sorted(expected_params) == sorted(actual_params) -def test_list_memories_rest_unset_required_fields(): - transport = transports.MemoryBankServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_memories._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_memories_rest_flattened(): client = MemoryBankServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4226,16 +4054,18 @@ def test_delete_memory_rest_required_fields(request_type=memory_bank_service.Del # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_memory._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteMemory, + "_BaseDeleteMemory__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_memory._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4281,13 +4111,6 @@ def test_delete_memory_rest_required_fields(request_type=memory_bank_service.Del assert sorted(expected_params) == sorted(actual_params) -def test_delete_memory_rest_unset_required_fields(): - transport = transports.MemoryBankServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_memory._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_memory_rest_flattened(): client = MemoryBankServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4392,16 +4215,18 @@ def test_generate_memories_rest_required_fields(request_type=memory_bank_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_memories._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGenerateMemories, + "_BaseGenerateMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_memories._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -4448,13 +4273,6 @@ def test_generate_memories_rest_required_fields(request_type=memory_bank_service assert sorted(expected_params) == sorted(actual_params) -def test_generate_memories_rest_unset_required_fields(): - transport = transports.MemoryBankServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.generate_memories._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_generate_memories_rest_flattened(): client = MemoryBankServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4555,16 +4373,18 @@ def test_retrieve_memories_rest_required_fields(request_type=memory_bank_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).retrieve_memories._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRetrieveMemories, + "_BaseRetrieveMemories__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).retrieve_memories._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -4614,13 +4434,6 @@ def test_retrieve_memories_rest_required_fields(request_type=memory_bank_service assert sorted(expected_params) == sorted(actual_params) -def test_retrieve_memories_rest_unset_required_fields(): - transport = transports.MemoryBankServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.retrieve_memories._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "scope", ))) - - def test_retrieve_memories_rest_flattened(): client = MemoryBankServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4991,6 +4804,8 @@ async def test_get_memory_empty_call_grpc_asyncio(): display_name='display_name_value', description='description_value', fact='fact_value', + memory_type=memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION, + context='context_value', )) await client.get_memory(request=None) @@ -5172,7 +4987,7 @@ def test_create_memory_rest_call_success(request_type): # send a request that will satisfy transcoding request_init = {'parent': 'projects/sample1/locations/sample2/reasoningEngines/sample3'} - request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'name': 'name_value', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}} + request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'revision_expire_time': {}, 'revision_ttl': {}, 'disable_memory_revisions': True, 'name': 'name_value', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}, 'topics': [{'custom_memory_topic_label': 'custom_memory_topic_label_value', 'managed_memory_topic': 1}], 'revision_labels': {}, 'metadata': {}, 'memory_type': 1, 'structured_content': {'data': {'fields': {}}, 'schema_id': 'schema_id_value'}, 'context': 'context_value'} # The version of a generated dependency at test runtime may differ from the version used during generation. # Delete any fields which are not present in the current runtime dependency # See https://github.com/googleapis/gapic-generator-python/issues/1748 @@ -5347,6 +5162,9 @@ def test_get_memory_rest_call_success(request_type): display_name='display_name_value', description='description_value', fact='fact_value', + memory_type=memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION, + context='context_value', + disable_memory_revisions=True, ) # Wrap the value into a proper Response obj @@ -5367,6 +5185,8 @@ def test_get_memory_rest_call_success(request_type): assert response.display_name == 'display_name_value' assert response.description == 'description_value' assert response.fact == 'fact_value' + assert response.memory_type == memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION + assert response.context == 'context_value' @pytest.mark.parametrize("null_interceptor", [True, False]) @@ -5449,7 +5269,7 @@ def test_update_memory_rest_call_success(request_type): # send a request that will satisfy transcoding request_init = {'memory': {'name': 'projects/sample1/locations/sample2/reasoningEngines/sample3/memories/sample4'}} - request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'name': 'projects/sample1/locations/sample2/reasoningEngines/sample3/memories/sample4', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}} + request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'revision_expire_time': {}, 'revision_ttl': {}, 'disable_memory_revisions': True, 'name': 'projects/sample1/locations/sample2/reasoningEngines/sample3/memories/sample4', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}, 'topics': [{'custom_memory_topic_label': 'custom_memory_topic_label_value', 'managed_memory_topic': 1}], 'revision_labels': {}, 'metadata': {}, 'memory_type': 1, 'structured_content': {'data': {'fields': {}}, 'schema_id': 'schema_id_value'}, 'context': 'context_value'} # The version of a generated dependency at test runtime may differ from the version used during generation. # Delete any fields which are not present in the current runtime dependency # See https://github.com/googleapis/gapic-generator-python/issues/1748 @@ -6744,7 +6564,7 @@ async def test_create_memory_rest_asyncio_call_success(request_type): # send a request that will satisfy transcoding request_init = {'parent': 'projects/sample1/locations/sample2/reasoningEngines/sample3'} - request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'name': 'name_value', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}} + request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'revision_expire_time': {}, 'revision_ttl': {}, 'disable_memory_revisions': True, 'name': 'name_value', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}, 'topics': [{'custom_memory_topic_label': 'custom_memory_topic_label_value', 'managed_memory_topic': 1}], 'revision_labels': {}, 'metadata': {}, 'memory_type': 1, 'structured_content': {'data': {'fields': {}}, 'schema_id': 'schema_id_value'}, 'context': 'context_value'} # The version of a generated dependency at test runtime may differ from the version used during generation. # Delete any fields which are not present in the current runtime dependency # See https://github.com/googleapis/gapic-generator-python/issues/1748 @@ -6926,6 +6746,9 @@ async def test_get_memory_rest_asyncio_call_success(request_type): display_name='display_name_value', description='description_value', fact='fact_value', + memory_type=memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION, + context='context_value', + disable_memory_revisions=True, ) # Wrap the value into a proper Response obj @@ -6946,6 +6769,8 @@ async def test_get_memory_rest_asyncio_call_success(request_type): assert response.display_name == 'display_name_value' assert response.description == 'description_value' assert response.fact == 'fact_value' + assert response.memory_type == memory_bank.MemoryType.NATURAL_LANGUAGE_COLLECTION + assert response.context == 'context_value' @pytest.mark.asyncio @@ -7035,7 +6860,7 @@ async def test_update_memory_rest_asyncio_call_success(request_type): # send a request that will satisfy transcoding request_init = {'memory': {'name': 'projects/sample1/locations/sample2/reasoningEngines/sample3/memories/sample4'}} - request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'name': 'projects/sample1/locations/sample2/reasoningEngines/sample3/memories/sample4', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}} + request_init["memory"] = {'expire_time': {'seconds': 751, 'nanos': 543}, 'ttl': {'seconds': 751, 'nanos': 543}, 'revision_expire_time': {}, 'revision_ttl': {}, 'disable_memory_revisions': True, 'name': 'projects/sample1/locations/sample2/reasoningEngines/sample3/memories/sample4', 'display_name': 'display_name_value', 'description': 'description_value', 'create_time': {}, 'update_time': {}, 'fact': 'fact_value', 'scope': {}, 'topics': [{'custom_memory_topic_label': 'custom_memory_topic_label_value', 'managed_memory_topic': 1}], 'revision_labels': {}, 'metadata': {}, 'memory_type': 1, 'structured_content': {'data': {'fields': {}}, 'schema_id': 'schema_id_value'}, 'context': 'context_value'} # The version of a generated dependency at test runtime may differ from the version used during generation. # Delete any fields which are not present in the current runtime dependency # See https://github.com/googleapis/gapic-generator-python/issues/1748 diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_metadata_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_metadata_service.py index 275cb85651..3102817435 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_metadata_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_metadata_service.py @@ -101,6 +101,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -144,150 +156,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert MetadataServiceClient._get_default_mtls_endpoint(None) is None - assert MetadataServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert MetadataServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert MetadataServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MetadataServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - MetadataServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert MetadataServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert MetadataServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert MetadataServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - MetadataServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert MetadataServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert MetadataServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - MetadataServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert MetadataServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert MetadataServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -301,40 +169,6 @@ def test__get_client_cert_source(): assert MetadataServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert MetadataServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(MetadataServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MetadataServiceClient)) -@mock.patch.object(MetadataServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MetadataServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = MetadataServiceClient._DEFAULT_UNIVERSE - default_endpoint = MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = MetadataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert MetadataServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert MetadataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - assert MetadataServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert MetadataServiceClient._get_api_endpoint(None, None, default_universe, "always") == MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - assert MetadataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == MetadataServiceClient.DEFAULT_MTLS_ENDPOINT - assert MetadataServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert MetadataServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - MetadataServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert MetadataServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert MetadataServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert MetadataServiceClient._get_universe_domain(None, None) == MetadataServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - MetadataServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -719,11 +553,12 @@ def test_metadata_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -766,11 +601,12 @@ def test_metadata_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -12460,17 +12296,20 @@ def test_create_metadata_store_rest_required_fields(request_type=metadata_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateMetadataStore, + "_BaseCreateMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("metadata_store_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("metadataStoreId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12518,13 +12357,6 @@ def test_create_metadata_store_rest_required_fields(request_type=metadata_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_metadata_store_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_metadata_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("metadataStoreId", )) & set(("parent", "metadataStore", ))) - - def test_create_metadata_store_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12629,16 +12461,18 @@ def test_get_metadata_store_rest_required_fields(request_type=metadata_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetMetadataStore, + "_BaseGetMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_store._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -12687,13 +12521,6 @@ def test_get_metadata_store_rest_required_fields(request_type=metadata_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_metadata_store_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_metadata_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_metadata_store_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12796,17 +12623,20 @@ def test_list_metadata_stores_rest_required_fields(request_type=metadata_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_stores._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListMetadataStores, + "_BaseListMetadataStores__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_stores._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12856,13 +12686,6 @@ def test_list_metadata_stores_rest_required_fields(request_type=metadata_service assert sorted(expected_params) == sorted(actual_params) -def test_list_metadata_stores_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_metadata_stores._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_metadata_stores_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13034,17 +12857,20 @@ def test_delete_metadata_store_rest_required_fields(request_type=metadata_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_metadata_store._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteMetadataStore, + "_BaseDeleteMetadataStore__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_metadata_store._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -13091,13 +12917,6 @@ def test_delete_metadata_store_rest_required_fields(request_type=metadata_servic assert sorted(expected_params) == sorted(actual_params) -def test_delete_metadata_store_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_metadata_store._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_metadata_store_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13198,17 +13017,20 @@ def test_create_artifact_rest_required_fields(request_type=metadata_service.Crea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateArtifact, + "_BaseCreateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_artifact._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("artifact_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("artifactId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13259,13 +13081,6 @@ def test_create_artifact_rest_required_fields(request_type=metadata_service.Crea assert sorted(expected_params) == sorted(actual_params) -def test_create_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(("artifactId", )) & set(("parent", "artifact", ))) - - def test_create_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13372,16 +13187,18 @@ def test_get_artifact_rest_required_fields(request_type=metadata_service.GetArti # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetArtifact, + "_BaseGetArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_artifact._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13430,13 +13247,6 @@ def test_get_artifact_rest_required_fields(request_type=metadata_service.GetArti assert sorted(expected_params) == sorted(actual_params) -def test_get_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13539,17 +13349,20 @@ def test_list_artifacts_rest_required_fields(request_type=metadata_service.ListA # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_artifacts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListArtifacts, + "_BaseListArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_artifacts._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13599,13 +13412,6 @@ def test_list_artifacts_rest_required_fields(request_type=metadata_service.ListA assert sorted(expected_params) == sorted(actual_params) -def test_list_artifacts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_artifacts._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_artifacts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13772,15 +13578,18 @@ def test_update_artifact_rest_required_fields(request_type=metadata_service.Upda # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateArtifact, + "_BaseUpdateArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_artifact._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("allow_missing", "update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("allowMissing", "updateMask", )) # verify required fields with non-default values are left alone @@ -13829,13 +13638,6 @@ def test_update_artifact_rest_required_fields(request_type=metadata_service.Upda assert sorted(expected_params) == sorted(actual_params) -def test_update_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(("allowMissing", "updateMask", )) & set(("artifact", ))) - - def test_update_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13944,17 +13746,20 @@ def test_delete_artifact_rest_required_fields(request_type=metadata_service.Dele # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_artifact._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteArtifact, + "_BaseDeleteArtifact__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_artifact._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("etag", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -14001,13 +13806,6 @@ def test_delete_artifact_rest_required_fields(request_type=metadata_service.Dele assert sorted(expected_params) == sorted(actual_params) -def test_delete_artifact_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_artifact._get_unset_required_fields({}) - assert set(unset_fields) == (set(("etag", )) & set(("name", ))) - - def test_delete_artifact_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14113,7 +13911,12 @@ def test_purge_artifacts_rest_required_fields(request_type=metadata_service.Purg # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_artifacts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePurgeArtifacts, + "_BasePurgeArtifacts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -14121,9 +13924,6 @@ def test_purge_artifacts_rest_required_fields(request_type=metadata_service.Purg jsonified_request["parent"] = 'parent_value' jsonified_request["filter"] = 'filter_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_artifacts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14172,13 +13972,6 @@ def test_purge_artifacts_rest_required_fields(request_type=metadata_service.Purg assert sorted(expected_params) == sorted(actual_params) -def test_purge_artifacts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.purge_artifacts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "filter", ))) - - def test_purge_artifacts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14279,17 +14072,20 @@ def test_create_context_rest_required_fields(request_type=metadata_service.Creat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateContext, + "_BaseCreateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_context._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("context_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("contextId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14340,13 +14136,6 @@ def test_create_context_rest_required_fields(request_type=metadata_service.Creat assert sorted(expected_params) == sorted(actual_params) -def test_create_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(("contextId", )) & set(("parent", "context", ))) - - def test_create_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14453,16 +14242,18 @@ def test_get_context_rest_required_fields(request_type=metadata_service.GetConte # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetContext, + "_BaseGetContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_context._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14511,13 +14302,6 @@ def test_get_context_rest_required_fields(request_type=metadata_service.GetConte assert sorted(expected_params) == sorted(actual_params) -def test_get_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14620,17 +14404,20 @@ def test_list_contexts_rest_required_fields(request_type=metadata_service.ListCo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListContexts, + "_BaseListContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_contexts._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14680,13 +14467,6 @@ def test_list_contexts_rest_required_fields(request_type=metadata_service.ListCo assert sorted(expected_params) == sorted(actual_params) -def test_list_contexts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_contexts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14853,15 +14633,18 @@ def test_update_context_rest_required_fields(request_type=metadata_service.Updat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateContext, + "_BaseUpdateContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_context._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("allow_missing", "update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("allowMissing", "updateMask", )) # verify required fields with non-default values are left alone @@ -14910,13 +14693,6 @@ def test_update_context_rest_required_fields(request_type=metadata_service.Updat assert sorted(expected_params) == sorted(actual_params) -def test_update_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(("allowMissing", "updateMask", )) & set(("context", ))) - - def test_update_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15025,17 +14801,20 @@ def test_delete_context_rest_required_fields(request_type=metadata_service.Delet # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_context._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteContext, + "_BaseDeleteContext__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_context._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("etag", "force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -15082,13 +14861,6 @@ def test_delete_context_rest_required_fields(request_type=metadata_service.Delet assert sorted(expected_params) == sorted(actual_params) -def test_delete_context_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_context._get_unset_required_fields({}) - assert set(unset_fields) == (set(("etag", "force", )) & set(("name", ))) - - def test_delete_context_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15194,7 +14966,12 @@ def test_purge_contexts_rest_required_fields(request_type=metadata_service.Purge # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePurgeContexts, + "_BasePurgeContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -15202,9 +14979,6 @@ def test_purge_contexts_rest_required_fields(request_type=metadata_service.Purge jsonified_request["parent"] = 'parent_value' jsonified_request["filter"] = 'filter_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -15253,13 +15027,6 @@ def test_purge_contexts_rest_required_fields(request_type=metadata_service.Purge assert sorted(expected_params) == sorted(actual_params) -def test_purge_contexts_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.purge_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "filter", ))) - - def test_purge_contexts_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15360,16 +15127,18 @@ def test_add_context_artifacts_and_executions_rest_required_fields(request_type= # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_artifacts_and_executions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddContextArtifactsAndExecutions, + "_BaseAddContextArtifactsAndExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_artifacts_and_executions._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15419,13 +15188,6 @@ def test_add_context_artifacts_and_executions_rest_required_fields(request_type= assert sorted(expected_params) == sorted(actual_params) -def test_add_context_artifacts_and_executions_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_context_artifacts_and_executions._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_add_context_artifacts_and_executions_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15532,16 +15294,18 @@ def test_add_context_children_rest_required_fields(request_type=metadata_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_children._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddContextChildren, + "_BaseAddContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_context_children._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15591,13 +15355,6 @@ def test_add_context_children_rest_required_fields(request_type=metadata_service assert sorted(expected_params) == sorted(actual_params) -def test_add_context_children_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_context_children._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_add_context_children_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15702,16 +15459,18 @@ def test_remove_context_children_rest_required_fields(request_type=metadata_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_context_children._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRemoveContextChildren, + "_BaseRemoveContextChildren__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).remove_context_children._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15761,13 +15520,6 @@ def test_remove_context_children_rest_required_fields(request_type=metadata_serv assert sorted(expected_params) == sorted(actual_params) -def test_remove_context_children_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.remove_context_children._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_remove_context_children_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15872,16 +15624,18 @@ def test_query_context_lineage_subgraph_rest_required_fields(request_type=metada # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_context_lineage_subgraph._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryContextLineageSubgraph, + "_BaseQueryContextLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["context"] = 'context_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_context_lineage_subgraph._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "context" in jsonified_request assert jsonified_request["context"] == 'context_value' @@ -15930,13 +15684,6 @@ def test_query_context_lineage_subgraph_rest_required_fields(request_type=metada assert sorted(expected_params) == sorted(actual_params) -def test_query_context_lineage_subgraph_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_context_lineage_subgraph._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("context", ))) - - def test_query_context_lineage_subgraph_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16039,17 +15786,20 @@ def test_create_execution_rest_required_fields(request_type=metadata_service.Cre # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateExecution, + "_BaseCreateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_execution._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("execution_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("executionId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -16100,13 +15850,6 @@ def test_create_execution_rest_required_fields(request_type=metadata_service.Cre assert sorted(expected_params) == sorted(actual_params) -def test_create_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(("executionId", )) & set(("parent", "execution", ))) - - def test_create_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16213,16 +15956,18 @@ def test_get_execution_rest_required_fields(request_type=metadata_service.GetExe # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetExecution, + "_BaseGetExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_execution._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -16271,13 +16016,6 @@ def test_get_execution_rest_required_fields(request_type=metadata_service.GetExe assert sorted(expected_params) == sorted(actual_params) -def test_get_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16380,17 +16118,20 @@ def test_list_executions_rest_required_fields(request_type=metadata_service.List # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_executions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListExecutions, + "_BaseListExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_executions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -16440,13 +16181,6 @@ def test_list_executions_rest_required_fields(request_type=metadata_service.List assert sorted(expected_params) == sorted(actual_params) -def test_list_executions_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_executions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_executions_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16613,15 +16347,18 @@ def test_update_execution_rest_required_fields(request_type=metadata_service.Upd # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateExecution, + "_BaseUpdateExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_execution._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("allow_missing", "update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("allowMissing", "updateMask", )) # verify required fields with non-default values are left alone @@ -16670,13 +16407,6 @@ def test_update_execution_rest_required_fields(request_type=metadata_service.Upd assert sorted(expected_params) == sorted(actual_params) -def test_update_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(("allowMissing", "updateMask", )) & set(("execution", ))) - - def test_update_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16785,17 +16515,20 @@ def test_delete_execution_rest_required_fields(request_type=metadata_service.Del # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_execution._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteExecution, + "_BaseDeleteExecution__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_execution._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("etag", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -16842,13 +16575,6 @@ def test_delete_execution_rest_required_fields(request_type=metadata_service.Del assert sorted(expected_params) == sorted(actual_params) -def test_delete_execution_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_execution._get_unset_required_fields({}) - assert set(unset_fields) == (set(("etag", )) & set(("name", ))) - - def test_delete_execution_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16954,7 +16680,12 @@ def test_purge_executions_rest_required_fields(request_type=metadata_service.Pur # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_executions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePurgeExecutions, + "_BasePurgeExecutions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -16962,9 +16693,6 @@ def test_purge_executions_rest_required_fields(request_type=metadata_service.Pur jsonified_request["parent"] = 'parent_value' jsonified_request["filter"] = 'filter_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).purge_executions._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -17013,13 +16741,6 @@ def test_purge_executions_rest_required_fields(request_type=metadata_service.Pur assert sorted(expected_params) == sorted(actual_params) -def test_purge_executions_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.purge_executions._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "filter", ))) - - def test_purge_executions_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17120,16 +16841,18 @@ def test_add_execution_events_rest_required_fields(request_type=metadata_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_execution_events._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddExecutionEvents, + "_BaseAddExecutionEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["execution"] = 'execution_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_execution_events._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "execution" in jsonified_request assert jsonified_request["execution"] == 'execution_value' @@ -17179,13 +16902,6 @@ def test_add_execution_events_rest_required_fields(request_type=metadata_service assert sorted(expected_params) == sorted(actual_params) -def test_add_execution_events_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_execution_events._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("execution", ))) - - def test_add_execution_events_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17290,16 +17006,18 @@ def test_query_execution_inputs_and_outputs_rest_required_fields(request_type=me # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_execution_inputs_and_outputs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryExecutionInputsAndOutputs, + "_BaseQueryExecutionInputsAndOutputs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["execution"] = 'execution_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_execution_inputs_and_outputs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "execution" in jsonified_request assert jsonified_request["execution"] == 'execution_value' @@ -17348,13 +17066,6 @@ def test_query_execution_inputs_and_outputs_rest_required_fields(request_type=me assert sorted(expected_params) == sorted(actual_params) -def test_query_execution_inputs_and_outputs_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_execution_inputs_and_outputs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("execution", ))) - - def test_query_execution_inputs_and_outputs_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17457,17 +17168,20 @@ def test_create_metadata_schema_rest_required_fields(request_type=metadata_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_schema._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateMetadataSchema, + "_BaseCreateMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_metadata_schema._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("metadata_schema_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("metadataSchemaId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17518,13 +17232,6 @@ def test_create_metadata_schema_rest_required_fields(request_type=metadata_servi assert sorted(expected_params) == sorted(actual_params) -def test_create_metadata_schema_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_metadata_schema._get_unset_required_fields({}) - assert set(unset_fields) == (set(("metadataSchemaId", )) & set(("parent", "metadataSchema", ))) - - def test_create_metadata_schema_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17631,16 +17338,18 @@ def test_get_metadata_schema_rest_required_fields(request_type=metadata_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_schema._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetMetadataSchema, + "_BaseGetMetadataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_metadata_schema._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -17689,13 +17398,6 @@ def test_get_metadata_schema_rest_required_fields(request_type=metadata_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_metadata_schema_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_metadata_schema._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_metadata_schema_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -17798,17 +17500,20 @@ def test_list_metadata_schemas_rest_required_fields(request_type=metadata_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_schemas._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListMetadataSchemas, + "_BaseListMetadataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_metadata_schemas._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -17858,13 +17563,6 @@ def test_list_metadata_schemas_rest_required_fields(request_type=metadata_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_metadata_schemas_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_metadata_schemas._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_metadata_schemas_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -18032,17 +17730,20 @@ def test_query_artifact_lineage_subgraph_rest_required_fields(request_type=metad # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_artifact_lineage_subgraph._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryArtifactLineageSubgraph, + "_BaseQueryArtifactLineageSubgraph__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["artifact"] = 'artifact_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_artifact_lineage_subgraph._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "max_hops", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "maxHops", )) # verify required fields with non-default values are left alone assert "artifact" in jsonified_request @@ -18092,13 +17793,6 @@ def test_query_artifact_lineage_subgraph_rest_required_fields(request_type=metad assert sorted(expected_params) == sorted(actual_params) -def test_query_artifact_lineage_subgraph_rest_unset_required_fields(): - transport = transports.MetadataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_artifact_lineage_subgraph._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "maxHops", )) & set(("artifact", ))) - - def test_query_artifact_lineage_subgraph_rest_flattened(): client = MetadataServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_migration_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_migration_service.py index 973fbffad4..6d1a328f6f 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_migration_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_migration_service.py @@ -84,6 +84,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -127,150 +139,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert MigrationServiceClient._get_default_mtls_endpoint(None) is None - assert MigrationServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert MigrationServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert MigrationServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MigrationServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - MigrationServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert MigrationServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert MigrationServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert MigrationServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - MigrationServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert MigrationServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert MigrationServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - MigrationServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert MigrationServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert MigrationServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -284,40 +152,6 @@ def test__get_client_cert_source(): assert MigrationServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert MigrationServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(MigrationServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MigrationServiceClient)) -@mock.patch.object(MigrationServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(MigrationServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = MigrationServiceClient._DEFAULT_UNIVERSE - default_endpoint = MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = MigrationServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert MigrationServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert MigrationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - assert MigrationServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert MigrationServiceClient._get_api_endpoint(None, None, default_universe, "always") == MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - assert MigrationServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == MigrationServiceClient.DEFAULT_MTLS_ENDPOINT - assert MigrationServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert MigrationServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - MigrationServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert MigrationServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert MigrationServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert MigrationServiceClient._get_universe_domain(None, None) == MigrationServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - MigrationServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -702,11 +536,12 @@ def test_migration_service_client_get_mtls_endpoint_and_cert_source(client_class for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -749,11 +584,12 @@ def test_migration_service_client_get_mtls_endpoint_and_cert_source(client_class for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1873,16 +1709,18 @@ def test_search_migratable_resources_rest_required_fields(request_type=migration # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_migratable_resources._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchMigratableResources, + "_BaseSearchMigratableResources__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_migratable_resources._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -1932,13 +1770,6 @@ def test_search_migratable_resources_rest_required_fields(request_type=migration assert sorted(expected_params) == sorted(actual_params) -def test_search_migratable_resources_rest_unset_required_fields(): - transport = transports.MigrationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_migratable_resources._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_search_migratable_resources_rest_flattened(): client = MigrationServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2110,16 +1941,18 @@ def test_batch_migrate_resources_rest_required_fields(request_type=migration_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_migrate_resources._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchMigrateResources, + "_BaseBatchMigrateResources__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_migrate_resources._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2166,13 +1999,6 @@ def test_batch_migrate_resources_rest_required_fields(request_type=migration_ser assert sorted(expected_params) == sorted(actual_params) -def test_batch_migrate_resources_rest_unset_required_fields(): - transport = transports.MigrationServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_migrate_resources._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "migrateResourceRequests", ))) - - def test_batch_migrate_resources_rest_flattened(): client = MigrationServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_model_garden_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_model_garden_service.py index e6ad4627dc..6be1f68cab 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_model_garden_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_model_garden_service.py @@ -92,6 +92,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -135,150 +147,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ModelGardenServiceClient._get_default_mtls_endpoint(None) is None - assert ModelGardenServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ModelGardenServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ModelGardenServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelGardenServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ModelGardenServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ModelGardenServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelGardenServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ModelGardenServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ModelGardenServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ModelGardenServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ModelGardenServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -292,40 +160,6 @@ def test__get_client_cert_source(): assert ModelGardenServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ModelGardenServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ModelGardenServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelGardenServiceClient)) -@mock.patch.object(ModelGardenServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelGardenServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ModelGardenServiceClient._DEFAULT_UNIVERSE - default_endpoint = ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ModelGardenServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ModelGardenServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ModelGardenServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelGardenServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ModelGardenServiceClient._get_api_endpoint(None, None, default_universe, "always") == ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelGardenServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ModelGardenServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelGardenServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ModelGardenServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelGardenServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ModelGardenServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ModelGardenServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ModelGardenServiceClient._get_universe_domain(None, None) == ModelGardenServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ModelGardenServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -710,11 +544,12 @@ def test_model_garden_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -757,11 +592,12 @@ def test_model_garden_service_client_get_mtls_endpoint_and_cert_source(client_cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3268,17 +3104,20 @@ def test_get_publisher_model_rest_required_fields(request_type=model_garden_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_publisher_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetPublisherModel, + "_BaseGetPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_publisher_model._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("hugging_face_token", "include_equivalent_model_garden_model_deployment_configs", "is_hugging_face_model", "language_code", "view", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("huggingFaceToken", "includeEquivalentModelGardenModelDeploymentConfigs", "isHuggingFaceModel", "languageCode", "view", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -3328,13 +3167,6 @@ def test_get_publisher_model_rest_required_fields(request_type=model_garden_serv assert sorted(expected_params) == sorted(actual_params) -def test_get_publisher_model_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_publisher_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(("huggingFaceToken", "includeEquivalentModelGardenModelDeploymentConfigs", "isHuggingFaceModel", "languageCode", "view", )) & set(("name", ))) - - def test_get_publisher_model_rest_flattened(): client = ModelGardenServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3437,17 +3269,20 @@ def test_list_publisher_models_rest_required_fields(request_type=model_garden_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_publisher_models._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListPublisherModels, + "_BaseListPublisherModels__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_publisher_models._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "language_code", "list_all_versions", "order_by", "page_size", "page_token", "view", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "languageCode", "listAllVersions", "orderBy", "pageSize", "pageToken", "view", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3497,13 +3332,6 @@ def test_list_publisher_models_rest_required_fields(request_type=model_garden_se assert sorted(expected_params) == sorted(actual_params) -def test_list_publisher_models_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_publisher_models._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "languageCode", "listAllVersions", "orderBy", "pageSize", "pageToken", "view", )) & set(("parent", ))) - - def test_list_publisher_models_rest_flattened(): client = ModelGardenServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3675,16 +3503,18 @@ def test_deploy_rest_required_fields(request_type=model_garden_service.DeployReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeploy, + "_BaseDeploy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["destination"] = 'destination_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "destination" in jsonified_request assert jsonified_request["destination"] == 'destination_value' @@ -3731,13 +3561,6 @@ def test_deploy_rest_required_fields(request_type=model_garden_service.DeployReq assert sorted(expected_params) == sorted(actual_params) -def test_deploy_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.deploy._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("destination", ))) - - def test_deploy_publisher_model_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -3791,7 +3614,12 @@ def test_deploy_publisher_model_rest_required_fields(request_type=model_garden_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_publisher_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeployPublisherModel, + "_BaseDeployPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -3799,9 +3627,6 @@ def test_deploy_publisher_model_rest_required_fields(request_type=model_garden_s jsonified_request["model"] = 'model_value' jsonified_request["destination"] = 'destination_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).deploy_publisher_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "model" in jsonified_request assert jsonified_request["model"] == 'model_value' @@ -3850,13 +3675,6 @@ def test_deploy_publisher_model_rest_required_fields(request_type=model_garden_s assert sorted(expected_params) == sorted(actual_params) -def test_deploy_publisher_model_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.deploy_publisher_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("model", "destination", ))) - - def test_export_publisher_model_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -3910,7 +3728,12 @@ def test_export_publisher_model_rest_required_fields(request_type=model_garden_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_publisher_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportPublisherModel, + "_BaseExportPublisherModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -3918,9 +3741,6 @@ def test_export_publisher_model_rest_required_fields(request_type=model_garden_s jsonified_request["name"] = 'name_value' jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_publisher_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3969,13 +3789,6 @@ def test_export_publisher_model_rest_required_fields(request_type=model_garden_s assert sorted(expected_params) == sorted(actual_params) -def test_export_publisher_model_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_publisher_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "destination", "parent", ))) - - def test_check_publisher_model_eula_acceptance_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -4025,7 +3838,12 @@ def test_check_publisher_model_eula_acceptance_rest_required_fields(request_type # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).check_publisher_model_eula_acceptance._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCheckPublisherModelEulaAcceptance, + "_BaseCheckPublisherModelEulaAcceptance__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -4033,9 +3851,6 @@ def test_check_publisher_model_eula_acceptance_rest_required_fields(request_type jsonified_request["parent"] = 'parent_value' jsonified_request["publisherModel"] = 'publisher_model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).check_publisher_model_eula_acceptance._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -4087,13 +3902,6 @@ def test_check_publisher_model_eula_acceptance_rest_required_fields(request_type assert sorted(expected_params) == sorted(actual_params) -def test_check_publisher_model_eula_acceptance_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.check_publisher_model_eula_acceptance._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "publisherModel", ))) - - def test_check_publisher_model_eula_acceptance_rest_flattened(): client = ModelGardenServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4199,7 +4007,12 @@ def test_accept_publisher_model_eula_rest_required_fields(request_type=model_gar # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).accept_publisher_model_eula._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAcceptPublisherModelEula, + "_BaseAcceptPublisherModelEula__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -4207,9 +4020,6 @@ def test_accept_publisher_model_eula_rest_required_fields(request_type=model_gar jsonified_request["parent"] = 'parent_value' jsonified_request["publisherModel"] = 'publisher_model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).accept_publisher_model_eula._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -4261,13 +4071,6 @@ def test_accept_publisher_model_eula_rest_required_fields(request_type=model_gar assert sorted(expected_params) == sorted(actual_params) -def test_accept_publisher_model_eula_rest_unset_required_fields(): - transport = transports.ModelGardenServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.accept_publisher_model_eula._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "publisherModel", ))) - - def test_accept_publisher_model_eula_rest_flattened(): client = ModelGardenServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_model_monitoring_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_model_monitoring_service.py index f1abd5a266..f1f858e401 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_model_monitoring_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_model_monitoring_service.py @@ -107,6 +107,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -150,150 +162,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ModelMonitoringServiceClient._get_default_mtls_endpoint(None) is None - assert ModelMonitoringServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ModelMonitoringServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ModelMonitoringServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ModelMonitoringServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ModelMonitoringServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ModelMonitoringServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ModelMonitoringServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelMonitoringServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelMonitoringServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ModelMonitoringServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ModelMonitoringServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ModelMonitoringServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ModelMonitoringServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ModelMonitoringServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelMonitoringServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ModelMonitoringServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ModelMonitoringServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ModelMonitoringServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ModelMonitoringServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ModelMonitoringServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ModelMonitoringServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -307,40 +175,6 @@ def test__get_client_cert_source(): assert ModelMonitoringServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ModelMonitoringServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ModelMonitoringServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelMonitoringServiceClient)) -@mock.patch.object(ModelMonitoringServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelMonitoringServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ModelMonitoringServiceClient._DEFAULT_UNIVERSE - default_endpoint = ModelMonitoringServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ModelMonitoringServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ModelMonitoringServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ModelMonitoringServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ModelMonitoringServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelMonitoringServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ModelMonitoringServiceClient._get_api_endpoint(None, None, default_universe, "always") == ModelMonitoringServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelMonitoringServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ModelMonitoringServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelMonitoringServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ModelMonitoringServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelMonitoringServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ModelMonitoringServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ModelMonitoringServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ModelMonitoringServiceClient._get_universe_domain(None, None) == ModelMonitoringServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ModelMonitoringServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -725,11 +559,12 @@ def test_model_monitoring_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -772,11 +607,12 @@ def test_model_monitoring_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5379,17 +5215,20 @@ def test_create_model_monitor_rest_required_fields(request_type=model_monitoring # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateModelMonitor, + "_BaseCreateModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_monitor._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("model_monitor_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("modelMonitorId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5437,13 +5276,6 @@ def test_create_model_monitor_rest_required_fields(request_type=model_monitoring assert sorted(expected_params) == sorted(actual_params) -def test_create_model_monitor_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_model_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(("modelMonitorId", )) & set(("parent", "modelMonitor", ))) - - def test_create_model_monitor_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5549,15 +5381,18 @@ def test_update_model_monitor_rest_required_fields(request_type=model_monitoring # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateModelMonitor, + "_BaseUpdateModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model_monitor._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -5603,13 +5438,6 @@ def test_update_model_monitor_rest_required_fields(request_type=model_monitoring assert sorted(expected_params) == sorted(actual_params) -def test_update_model_monitor_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_model_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("modelMonitor", "updateMask", ))) - - def test_update_model_monitor_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5712,16 +5540,18 @@ def test_get_model_monitor_rest_required_fields(request_type=model_monitoring_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelMonitor, + "_BaseGetModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_monitor._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5770,13 +5600,6 @@ def test_get_model_monitor_rest_required_fields(request_type=model_monitoring_se assert sorted(expected_params) == sorted(actual_params) -def test_get_model_monitor_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_monitor_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5879,17 +5702,20 @@ def test_list_model_monitors_rest_required_fields(request_type=model_monitoring_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_monitors._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelMonitors, + "_BaseListModelMonitors__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_monitors._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5939,13 +5765,6 @@ def test_list_model_monitors_rest_required_fields(request_type=model_monitoring_ assert sorted(expected_params) == sorted(actual_params) -def test_list_model_monitors_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_monitors._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_monitors_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6117,17 +5936,20 @@ def test_delete_model_monitor_rest_required_fields(request_type=model_monitoring # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_monitor._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModelMonitor, + "_BaseDeleteModelMonitor__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_monitor._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -6174,13 +5996,6 @@ def test_delete_model_monitor_rest_required_fields(request_type=model_monitoring assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_monitor_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model_monitor._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_model_monitor_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6281,17 +6096,20 @@ def test_create_model_monitoring_job_rest_required_fields(request_type=model_mon # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateModelMonitoringJob, + "_BaseCreateModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_model_monitoring_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("model_monitoring_job_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("modelMonitoringJobId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6342,13 +6160,6 @@ def test_create_model_monitoring_job_rest_required_fields(request_type=model_mon assert sorted(expected_params) == sorted(actual_params) -def test_create_model_monitoring_job_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_model_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("modelMonitoringJobId", )) & set(("parent", "modelMonitoringJob", ))) - - def test_create_model_monitoring_job_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6453,16 +6264,18 @@ def test_get_model_monitoring_job_rest_required_fields(request_type=model_monito # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelMonitoringJob, + "_BaseGetModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6511,13 +6324,6 @@ def test_get_model_monitoring_job_rest_required_fields(request_type=model_monito assert sorted(expected_params) == sorted(actual_params) -def test_get_model_monitoring_job_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_monitoring_job_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6620,17 +6426,20 @@ def test_list_model_monitoring_jobs_rest_required_fields(request_type=model_moni # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_monitoring_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelMonitoringJobs, + "_BaseListModelMonitoringJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_monitoring_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6680,13 +6489,6 @@ def test_list_model_monitoring_jobs_rest_required_fields(request_type=model_moni assert sorted(expected_params) == sorted(actual_params) -def test_list_model_monitoring_jobs_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_monitoring_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_monitoring_jobs_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6858,16 +6660,18 @@ def test_delete_model_monitoring_job_rest_required_fields(request_type=model_mon # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_monitoring_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModelMonitoringJob, + "_BaseDeleteModelMonitoringJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_monitoring_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6913,13 +6717,6 @@ def test_delete_model_monitoring_job_rest_required_fields(request_type=model_mon assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_monitoring_job_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model_monitoring_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_model_monitoring_job_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7020,16 +6817,18 @@ def test_search_model_monitoring_stats_rest_required_fields(request_type=model_m # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_monitoring_stats._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchModelMonitoringStats, + "_BaseSearchModelMonitoringStats__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["modelMonitor"] = 'model_monitor_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_monitoring_stats._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "modelMonitor" in jsonified_request assert jsonified_request["modelMonitor"] == 'model_monitor_value' @@ -7079,13 +6878,6 @@ def test_search_model_monitoring_stats_rest_required_fields(request_type=model_m assert sorted(expected_params) == sorted(actual_params) -def test_search_model_monitoring_stats_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_model_monitoring_stats._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("modelMonitor", ))) - - def test_search_model_monitoring_stats_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7253,16 +7045,18 @@ def test_search_model_monitoring_alerts_rest_required_fields(request_type=model_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_monitoring_alerts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSearchModelMonitoringAlerts, + "_BaseSearchModelMonitoringAlerts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["modelMonitor"] = 'model_monitor_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).search_model_monitoring_alerts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "modelMonitor" in jsonified_request assert jsonified_request["modelMonitor"] == 'model_monitor_value' @@ -7312,13 +7106,6 @@ def test_search_model_monitoring_alerts_rest_required_fields(request_type=model_ assert sorted(expected_params) == sorted(actual_params) -def test_search_model_monitoring_alerts_rest_unset_required_fields(): - transport = transports.ModelMonitoringServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.search_model_monitoring_alerts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("modelMonitor", ))) - - def test_search_model_monitoring_alerts_rest_flattened(): client = ModelMonitoringServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_model_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_model_service.py index 06db6f6ced..46f8bd69e6 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_model_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_model_service.py @@ -102,6 +102,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -145,150 +157,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ModelServiceClient._get_default_mtls_endpoint(None) is None - assert ModelServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ModelServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ModelServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ModelServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ModelServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ModelServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ModelServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ModelServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ModelServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ModelServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ModelServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ModelServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ModelServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ModelServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -302,40 +170,6 @@ def test__get_client_cert_source(): assert ModelServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ModelServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ModelServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelServiceClient)) -@mock.patch.object(ModelServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ModelServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ModelServiceClient._DEFAULT_UNIVERSE - default_endpoint = ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ModelServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ModelServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ModelServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ModelServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ModelServiceClient._get_api_endpoint(None, None, default_universe, "always") == ModelServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ModelServiceClient.DEFAULT_MTLS_ENDPOINT - assert ModelServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ModelServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ModelServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ModelServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ModelServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ModelServiceClient._get_universe_domain(None, None) == ModelServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ModelServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -720,11 +554,12 @@ def test_model_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -767,11 +602,12 @@ def test_model_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -8531,16 +8367,18 @@ def test_upload_model_rest_required_fields(request_type=model_service.UploadMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUploadModel, + "_BaseUploadModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -8587,13 +8425,6 @@ def test_upload_model_rest_required_fields(request_type=model_service.UploadMode assert sorted(expected_params) == sorted(actual_params) -def test_upload_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upload_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "model", ))) - - def test_upload_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8696,16 +8527,18 @@ def test_get_model_rest_required_fields(request_type=model_service.GetModelReque # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModel, + "_BaseGetModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8754,13 +8587,6 @@ def test_get_model_rest_required_fields(request_type=model_service.GetModelReque assert sorted(expected_params) == sorted(actual_params) -def test_get_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8863,17 +8689,20 @@ def test_list_models_rest_required_fields(request_type=model_service.ListModelsR # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_models._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModels, + "_BaseListModels__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_models._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8923,13 +8752,6 @@ def test_list_models_rest_required_fields(request_type=model_service.ListModelsR assert sorted(expected_params) == sorted(actual_params) -def test_list_models_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_models._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_models_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9097,17 +8919,20 @@ def test_list_model_versions_rest_required_fields(request_type=model_service.Lis # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_versions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelVersions, + "_BaseListModelVersions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_versions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9157,13 +8982,6 @@ def test_list_model_versions_rest_required_fields(request_type=model_service.Lis assert sorted(expected_params) == sorted(actual_params) -def test_list_model_versions_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_versions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("name", ))) - - def test_list_model_versions_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9331,17 +9149,20 @@ def test_list_model_version_checkpoints_rest_required_fields(request_type=model_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_version_checkpoints._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelVersionCheckpoints, + "_BaseListModelVersionCheckpoints__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_version_checkpoints._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9391,13 +9212,6 @@ def test_list_model_version_checkpoints_rest_required_fields(request_type=model_ assert sorted(expected_params) == sorted(actual_params) -def test_list_model_version_checkpoints_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_version_checkpoints._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("name", ))) - - def test_list_model_version_checkpoints_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9564,15 +9378,18 @@ def test_update_model_rest_required_fields(request_type=model_service.UpdateMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateModel, + "_BaseUpdateModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_model._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -9621,13 +9438,6 @@ def test_update_model_rest_required_fields(request_type=model_service.UpdateMode assert sorted(expected_params) == sorted(actual_params) -def test_update_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("model", "updateMask", ))) - - def test_update_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9736,16 +9546,18 @@ def test_update_explanation_dataset_rest_required_fields(request_type=model_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_explanation_dataset._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateExplanationDataset, + "_BaseUpdateExplanationDataset__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["model"] = 'model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_explanation_dataset._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "model" in jsonified_request assert jsonified_request["model"] == 'model_value' @@ -9792,13 +9604,6 @@ def test_update_explanation_dataset_rest_required_fields(request_type=model_serv assert sorted(expected_params) == sorted(actual_params) -def test_update_explanation_dataset_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_explanation_dataset._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("model", ))) - - def test_update_explanation_dataset_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9903,16 +9708,18 @@ def test_delete_model_rest_required_fields(request_type=model_service.DeleteMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModel, + "_BaseDeleteModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9958,13 +9765,6 @@ def test_delete_model_rest_required_fields(request_type=model_service.DeleteMode assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10069,16 +9869,18 @@ def test_delete_model_version_rest_required_fields(request_type=model_service.De # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_version._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteModelVersion, + "_BaseDeleteModelVersion__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_model_version._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10124,13 +9926,6 @@ def test_delete_model_version_rest_required_fields(request_type=model_service.De assert sorted(expected_params) == sorted(actual_params) -def test_delete_model_version_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_model_version._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_model_version_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10232,7 +10027,12 @@ def test_merge_version_aliases_rest_required_fields(request_type=model_service.M # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).merge_version_aliases._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseMergeVersionAliases, + "_BaseMergeVersionAliases__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -10240,9 +10040,6 @@ def test_merge_version_aliases_rest_required_fields(request_type=model_service.M jsonified_request["name"] = 'name_value' jsonified_request["versionAliases"] = 'version_aliases_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).merge_version_aliases._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10294,13 +10091,6 @@ def test_merge_version_aliases_rest_required_fields(request_type=model_service.M assert sorted(expected_params) == sorted(actual_params) -def test_merge_version_aliases_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.merge_version_aliases._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "versionAliases", ))) - - def test_merge_version_aliases_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10409,16 +10199,18 @@ def test_export_model_rest_required_fields(request_type=model_service.ExportMode # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportModel, + "_BaseExportModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10465,13 +10257,6 @@ def test_export_model_rest_required_fields(request_type=model_service.ExportMode assert sorted(expected_params) == sorted(actual_params) -def test_export_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "outputConfig", ))) - - def test_export_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10579,7 +10364,12 @@ def test_copy_model_rest_required_fields(request_type=model_service.CopyModelReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).copy_model._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCopyModel, + "_BaseCopyModel__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -10587,9 +10377,6 @@ def test_copy_model_rest_required_fields(request_type=model_service.CopyModelReq jsonified_request["parent"] = 'parent_value' jsonified_request["sourceModel"] = 'source_model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).copy_model._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10638,13 +10425,6 @@ def test_copy_model_rest_required_fields(request_type=model_service.CopyModelReq assert sorted(expected_params) == sorted(actual_params) -def test_copy_model_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.copy_model._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "sourceModel", ))) - - def test_copy_model_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10747,16 +10527,18 @@ def test_import_model_evaluation_rest_required_fields(request_type=model_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_model_evaluation._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportModelEvaluation, + "_BaseImportModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_model_evaluation._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10806,13 +10588,6 @@ def test_import_model_evaluation_rest_required_fields(request_type=model_service assert sorted(expected_params) == sorted(actual_params) -def test_import_model_evaluation_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_model_evaluation._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "modelEvaluation", ))) - - def test_import_model_evaluation_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10917,16 +10692,18 @@ def test_batch_import_model_evaluation_slices_rest_required_fields(request_type= # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_model_evaluation_slices._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchImportModelEvaluationSlices, + "_BaseBatchImportModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_model_evaluation_slices._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10976,13 +10753,6 @@ def test_batch_import_model_evaluation_slices_rest_required_fields(request_type= assert sorted(expected_params) == sorted(actual_params) -def test_batch_import_model_evaluation_slices_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_import_model_evaluation_slices._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "modelEvaluationSlices", ))) - - def test_batch_import_model_evaluation_slices_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11087,16 +10857,18 @@ def test_batch_import_evaluated_annotations_rest_required_fields(request_type=mo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_evaluated_annotations._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchImportEvaluatedAnnotations, + "_BaseBatchImportEvaluatedAnnotations__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_import_evaluated_annotations._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -11146,13 +10918,6 @@ def test_batch_import_evaluated_annotations_rest_required_fields(request_type=mo assert sorted(expected_params) == sorted(actual_params) -def test_batch_import_evaluated_annotations_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_import_evaluated_annotations._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "evaluatedAnnotations", ))) - - def test_batch_import_evaluated_annotations_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11257,16 +11022,18 @@ def test_get_model_evaluation_rest_required_fields(request_type=model_service.Ge # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelEvaluation, + "_BaseGetModelEvaluation__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11315,13 +11082,6 @@ def test_get_model_evaluation_rest_required_fields(request_type=model_service.Ge assert sorted(expected_params) == sorted(actual_params) -def test_get_model_evaluation_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_evaluation._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_evaluation_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11424,17 +11184,20 @@ def test_list_model_evaluations_rest_required_fields(request_type=model_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluations._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelEvaluations, + "_BaseListModelEvaluations__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluations._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11484,13 +11247,6 @@ def test_list_model_evaluations_rest_required_fields(request_type=model_service. assert sorted(expected_params) == sorted(actual_params) -def test_list_model_evaluations_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_evaluations._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_evaluations_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11658,16 +11414,18 @@ def test_get_model_evaluation_slice_rest_required_fields(request_type=model_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation_slice._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetModelEvaluationSlice, + "_BaseGetModelEvaluationSlice__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_model_evaluation_slice._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11716,13 +11474,6 @@ def test_get_model_evaluation_slice_rest_required_fields(request_type=model_serv assert sorted(expected_params) == sorted(actual_params) -def test_get_model_evaluation_slice_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_model_evaluation_slice._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_model_evaluation_slice_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11825,17 +11576,20 @@ def test_list_model_evaluation_slices_rest_required_fields(request_type=model_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluation_slices._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListModelEvaluationSlices, + "_BaseListModelEvaluationSlices__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_model_evaluation_slices._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11885,13 +11639,6 @@ def test_list_model_evaluation_slices_rest_required_fields(request_type=model_se assert sorted(expected_params) == sorted(actual_params) -def test_list_model_evaluation_slices_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_model_evaluation_slices._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_model_evaluation_slices_rest_flattened(): client = ModelServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12060,7 +11807,12 @@ def test_recommend_spec_rest_required_fields(request_type=model_service.Recommen # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).recommend_spec._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRecommendSpec, + "_BaseRecommendSpec__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -12068,9 +11820,6 @@ def test_recommend_spec_rest_required_fields(request_type=model_service.Recommen jsonified_request["parent"] = 'parent_value' jsonified_request["gcsUri"] = 'gcs_uri_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).recommend_spec._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -12122,13 +11871,6 @@ def test_recommend_spec_rest_required_fields(request_type=model_service.Recommen assert sorted(expected_params) == sorted(actual_params) -def test_recommend_spec_rest_unset_required_fields(): - transport = transports.ModelServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.recommend_spec._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "gcsUri", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.ModelServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_notebook_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_notebook_service.py index 7191a85acb..64cb6fdc35 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_notebook_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_notebook_service.py @@ -105,6 +105,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -148,150 +160,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert NotebookServiceClient._get_default_mtls_endpoint(None) is None - assert NotebookServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert NotebookServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert NotebookServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert NotebookServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - NotebookServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert NotebookServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert NotebookServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert NotebookServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - NotebookServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert NotebookServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert NotebookServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - NotebookServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert NotebookServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert NotebookServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -305,40 +173,6 @@ def test__get_client_cert_source(): assert NotebookServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert NotebookServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(NotebookServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(NotebookServiceClient)) -@mock.patch.object(NotebookServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(NotebookServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = NotebookServiceClient._DEFAULT_UNIVERSE - default_endpoint = NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = NotebookServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert NotebookServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert NotebookServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - assert NotebookServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert NotebookServiceClient._get_api_endpoint(None, None, default_universe, "always") == NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - assert NotebookServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == NotebookServiceClient.DEFAULT_MTLS_ENDPOINT - assert NotebookServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert NotebookServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - NotebookServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert NotebookServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert NotebookServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert NotebookServiceClient._get_universe_domain(None, None) == NotebookServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - NotebookServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -723,11 +557,12 @@ def test_notebook_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -770,11 +605,12 @@ def test_notebook_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -6903,17 +6739,20 @@ def test_create_notebook_runtime_template_rest_required_fields(request_type=note # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateNotebookRuntimeTemplate, + "_BaseCreateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_runtime_template._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("notebook_runtime_template_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("notebookRuntimeTemplateId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6961,13 +6800,6 @@ def test_create_notebook_runtime_template_rest_required_fields(request_type=note assert sorted(expected_params) == sorted(actual_params) -def test_create_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(("notebookRuntimeTemplateId", )) & set(("parent", "notebookRuntimeTemplate", ))) - - def test_create_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7072,16 +6904,18 @@ def test_get_notebook_runtime_template_rest_required_fields(request_type=noteboo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNotebookRuntimeTemplate, + "_BaseGetNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime_template._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7130,13 +6964,6 @@ def test_get_notebook_runtime_template_rest_required_fields(request_type=noteboo assert sorted(expected_params) == sorted(actual_params) -def test_get_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7239,17 +7066,20 @@ def test_list_notebook_runtime_templates_rest_required_fields(request_type=noteb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtime_templates._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNotebookRuntimeTemplates, + "_BaseListNotebookRuntimeTemplates__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtime_templates._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -7299,13 +7129,6 @@ def test_list_notebook_runtime_templates_rest_required_fields(request_type=noteb assert sorted(expected_params) == sorted(actual_params) -def test_list_notebook_runtime_templates_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_notebook_runtime_templates._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_notebook_runtime_templates_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7477,16 +7300,18 @@ def test_delete_notebook_runtime_template_rest_required_fields(request_type=note # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNotebookRuntimeTemplate, + "_BaseDeleteNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime_template._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7532,13 +7357,6 @@ def test_delete_notebook_runtime_template_rest_required_fields(request_type=note assert sorted(expected_params) == sorted(actual_params) -def test_delete_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7638,15 +7456,18 @@ def test_update_notebook_runtime_template_rest_required_fields(request_type=note # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_notebook_runtime_template._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateNotebookRuntimeTemplate, + "_BaseUpdateNotebookRuntimeTemplate__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_notebook_runtime_template._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -7695,13 +7516,6 @@ def test_update_notebook_runtime_template_rest_required_fields(request_type=note assert sorted(expected_params) == sorted(actual_params) -def test_update_notebook_runtime_template_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_notebook_runtime_template._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("notebookRuntimeTemplate", "updateMask", ))) - - def test_update_notebook_runtime_template_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7811,7 +7625,12 @@ def test_assign_notebook_runtime_rest_required_fields(request_type=notebook_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assign_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAssignNotebookRuntime, + "_BaseAssignNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -7819,9 +7638,6 @@ def test_assign_notebook_runtime_rest_required_fields(request_type=notebook_serv jsonified_request["parent"] = 'parent_value' jsonified_request["notebookRuntimeTemplate"] = 'notebook_runtime_template_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).assign_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -7870,13 +7686,6 @@ def test_assign_notebook_runtime_rest_required_fields(request_type=notebook_serv assert sorted(expected_params) == sorted(actual_params) -def test_assign_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.assign_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "notebookRuntimeTemplate", "notebookRuntime", ))) - - def test_assign_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7983,16 +7792,18 @@ def test_get_notebook_runtime_rest_required_fields(request_type=notebook_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNotebookRuntime, + "_BaseGetNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8041,13 +7852,6 @@ def test_get_notebook_runtime_rest_required_fields(request_type=notebook_service assert sorted(expected_params) == sorted(actual_params) -def test_get_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8150,17 +7954,20 @@ def test_list_notebook_runtimes_rest_required_fields(request_type=notebook_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtimes._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNotebookRuntimes, + "_BaseListNotebookRuntimes__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_runtimes._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -8210,13 +8017,6 @@ def test_list_notebook_runtimes_rest_required_fields(request_type=notebook_servi assert sorted(expected_params) == sorted(actual_params) -def test_list_notebook_runtimes_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_notebook_runtimes._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_notebook_runtimes_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8388,16 +8188,18 @@ def test_delete_notebook_runtime_rest_required_fields(request_type=notebook_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNotebookRuntime, + "_BaseDeleteNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8443,13 +8245,6 @@ def test_delete_notebook_runtime_rest_required_fields(request_type=notebook_serv assert sorted(expected_params) == sorted(actual_params) -def test_delete_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8554,16 +8349,18 @@ def test_upgrade_notebook_runtime_rest_required_fields(request_type=notebook_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upgrade_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpgradeNotebookRuntime, + "_BaseUpgradeNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upgrade_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8610,13 +8407,6 @@ def test_upgrade_notebook_runtime_rest_required_fields(request_type=notebook_ser assert sorted(expected_params) == sorted(actual_params) -def test_upgrade_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upgrade_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_upgrade_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8721,16 +8511,18 @@ def test_start_notebook_runtime_rest_required_fields(request_type=notebook_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).start_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStartNotebookRuntime, + "_BaseStartNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).start_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8777,13 +8569,6 @@ def test_start_notebook_runtime_rest_required_fields(request_type=notebook_servi assert sorted(expected_params) == sorted(actual_params) -def test_start_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.start_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_start_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -8888,16 +8673,18 @@ def test_stop_notebook_runtime_rest_required_fields(request_type=notebook_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_notebook_runtime._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStopNotebookRuntime, + "_BaseStopNotebookRuntime__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_notebook_runtime._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -8944,13 +8731,6 @@ def test_stop_notebook_runtime_rest_required_fields(request_type=notebook_servic assert sorted(expected_params) == sorted(actual_params) -def test_stop_notebook_runtime_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stop_notebook_runtime._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_stop_notebook_runtime_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9055,17 +8835,20 @@ def test_create_notebook_execution_job_rest_required_fields(request_type=noteboo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_execution_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateNotebookExecutionJob, + "_BaseCreateNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_notebook_execution_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("notebook_execution_job_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("notebookExecutionJobId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9113,13 +8896,6 @@ def test_create_notebook_execution_job_rest_required_fields(request_type=noteboo assert sorted(expected_params) == sorted(actual_params) -def test_create_notebook_execution_job_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_notebook_execution_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("notebookExecutionJobId", )) & set(("parent", "notebookExecutionJob", ))) - - def test_create_notebook_execution_job_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9224,17 +9000,20 @@ def test_get_notebook_execution_job_rest_required_fields(request_type=notebook_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_execution_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetNotebookExecutionJob, + "_BaseGetNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_notebook_execution_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("view", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -9284,13 +9063,6 @@ def test_get_notebook_execution_job_rest_required_fields(request_type=notebook_s assert sorted(expected_params) == sorted(actual_params) -def test_get_notebook_execution_job_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_notebook_execution_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("view", )) & set(("name", ))) - - def test_get_notebook_execution_job_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9393,17 +9165,20 @@ def test_list_notebook_execution_jobs_rest_required_fields(request_type=notebook # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_execution_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListNotebookExecutionJobs, + "_BaseListNotebookExecutionJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_notebook_execution_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "view", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "view", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9453,13 +9228,6 @@ def test_list_notebook_execution_jobs_rest_required_fields(request_type=notebook assert sorted(expected_params) == sorted(actual_params) -def test_list_notebook_execution_jobs_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_notebook_execution_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "view", )) & set(("parent", ))) - - def test_list_notebook_execution_jobs_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9631,16 +9399,18 @@ def test_delete_notebook_execution_job_rest_required_fields(request_type=noteboo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_execution_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteNotebookExecutionJob, + "_BaseDeleteNotebookExecutionJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_notebook_execution_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9686,13 +9456,6 @@ def test_delete_notebook_execution_job_rest_required_fields(request_type=noteboo assert sorted(expected_params) == sorted(actual_params) -def test_delete_notebook_execution_job_rest_unset_required_fields(): - transport = transports.NotebookServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_notebook_execution_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_notebook_execution_job_rest_flattened(): client = NotebookServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_online_evaluator_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_online_evaluator_service.py index e50642d029..1c2db7cd0a 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_online_evaluator_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_online_evaluator_service.py @@ -92,6 +92,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -135,150 +147,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert OnlineEvaluatorServiceClient._get_default_mtls_endpoint(None) is None - assert OnlineEvaluatorServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert OnlineEvaluatorServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert OnlineEvaluatorServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert OnlineEvaluatorServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert OnlineEvaluatorServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert OnlineEvaluatorServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert OnlineEvaluatorServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert OnlineEvaluatorServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert OnlineEvaluatorServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - OnlineEvaluatorServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert OnlineEvaluatorServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert OnlineEvaluatorServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert OnlineEvaluatorServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert OnlineEvaluatorServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - OnlineEvaluatorServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert OnlineEvaluatorServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - OnlineEvaluatorServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert OnlineEvaluatorServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -292,40 +160,6 @@ def test__get_client_cert_source(): assert OnlineEvaluatorServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert OnlineEvaluatorServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(OnlineEvaluatorServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(OnlineEvaluatorServiceClient)) -@mock.patch.object(OnlineEvaluatorServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(OnlineEvaluatorServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = OnlineEvaluatorServiceClient._DEFAULT_UNIVERSE - default_endpoint = OnlineEvaluatorServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = OnlineEvaluatorServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert OnlineEvaluatorServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert OnlineEvaluatorServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == OnlineEvaluatorServiceClient.DEFAULT_MTLS_ENDPOINT - assert OnlineEvaluatorServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert OnlineEvaluatorServiceClient._get_api_endpoint(None, None, default_universe, "always") == OnlineEvaluatorServiceClient.DEFAULT_MTLS_ENDPOINT - assert OnlineEvaluatorServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == OnlineEvaluatorServiceClient.DEFAULT_MTLS_ENDPOINT - assert OnlineEvaluatorServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert OnlineEvaluatorServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - OnlineEvaluatorServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert OnlineEvaluatorServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert OnlineEvaluatorServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert OnlineEvaluatorServiceClient._get_universe_domain(None, None) == OnlineEvaluatorServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - OnlineEvaluatorServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -710,11 +544,12 @@ def test_online_evaluator_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -757,11 +592,12 @@ def test_online_evaluator_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3484,16 +3320,18 @@ def test_create_online_evaluator_rest_required_fields(request_type=online_evalua # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_online_evaluator._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateOnlineEvaluator, + "_BaseCreateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_online_evaluator._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3540,13 +3378,6 @@ def test_create_online_evaluator_rest_required_fields(request_type=online_evalua assert sorted(expected_params) == sorted(actual_params) -def test_create_online_evaluator_rest_unset_required_fields(): - transport = transports.OnlineEvaluatorServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_online_evaluator._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "onlineEvaluator", ))) - - def test_create_online_evaluator_rest_flattened(): client = OnlineEvaluatorServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3649,16 +3480,18 @@ def test_get_online_evaluator_rest_required_fields(request_type=online_evaluator # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_online_evaluator._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetOnlineEvaluator, + "_BaseGetOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_online_evaluator._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3707,13 +3540,6 @@ def test_get_online_evaluator_rest_required_fields(request_type=online_evaluator assert sorted(expected_params) == sorted(actual_params) -def test_get_online_evaluator_rest_unset_required_fields(): - transport = transports.OnlineEvaluatorServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_online_evaluator._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_online_evaluator_rest_flattened(): client = OnlineEvaluatorServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3819,15 +3645,18 @@ def test_update_online_evaluator_rest_required_fields(request_type=online_evalua # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_online_evaluator._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateOnlineEvaluator, + "_BaseUpdateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_online_evaluator._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3873,13 +3702,6 @@ def test_update_online_evaluator_rest_required_fields(request_type=online_evalua assert sorted(expected_params) == sorted(actual_params) -def test_update_online_evaluator_rest_unset_required_fields(): - transport = transports.OnlineEvaluatorServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_online_evaluator._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("onlineEvaluator", ))) - - def test_update_online_evaluator_rest_flattened(): client = OnlineEvaluatorServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3986,16 +3808,18 @@ def test_delete_online_evaluator_rest_required_fields(request_type=online_evalua # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_online_evaluator._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteOnlineEvaluator, + "_BaseDeleteOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_online_evaluator._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4041,13 +3865,6 @@ def test_delete_online_evaluator_rest_required_fields(request_type=online_evalua assert sorted(expected_params) == sorted(actual_params) -def test_delete_online_evaluator_rest_unset_required_fields(): - transport = transports.OnlineEvaluatorServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_online_evaluator._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_online_evaluator_rest_flattened(): client = OnlineEvaluatorServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4148,17 +3965,20 @@ def test_list_online_evaluators_rest_required_fields(request_type=online_evaluat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_online_evaluators._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListOnlineEvaluators, + "_BaseListOnlineEvaluators__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_online_evaluators._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4208,13 +4028,6 @@ def test_list_online_evaluators_rest_required_fields(request_type=online_evaluat assert sorted(expected_params) == sorted(actual_params) -def test_list_online_evaluators_rest_unset_required_fields(): - transport = transports.OnlineEvaluatorServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_online_evaluators._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_online_evaluators_rest_flattened(): client = OnlineEvaluatorServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4386,16 +4199,18 @@ def test_activate_online_evaluator_rest_required_fields(request_type=online_eval # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).activate_online_evaluator._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseActivateOnlineEvaluator, + "_BaseActivateOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).activate_online_evaluator._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4442,13 +4257,6 @@ def test_activate_online_evaluator_rest_required_fields(request_type=online_eval assert sorted(expected_params) == sorted(actual_params) -def test_activate_online_evaluator_rest_unset_required_fields(): - transport = transports.OnlineEvaluatorServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.activate_online_evaluator._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_activate_online_evaluator_rest_flattened(): client = OnlineEvaluatorServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4553,16 +4361,18 @@ def test_suspend_online_evaluator_rest_required_fields(request_type=online_evalu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).suspend_online_evaluator._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSuspendOnlineEvaluator, + "_BaseSuspendOnlineEvaluator__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).suspend_online_evaluator._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4609,13 +4419,6 @@ def test_suspend_online_evaluator_rest_required_fields(request_type=online_evalu assert sorted(expected_params) == sorted(actual_params) -def test_suspend_online_evaluator_rest_unset_required_fields(): - transport = transports.OnlineEvaluatorServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.suspend_online_evaluator._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_suspend_online_evaluator_rest_flattened(): client = OnlineEvaluatorServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_persistent_resource_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_persistent_resource_service.py index 4a62d84ef9..65f7b7b64b 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_persistent_resource_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_persistent_resource_service.py @@ -96,6 +96,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -139,150 +151,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert PersistentResourceServiceClient._get_default_mtls_endpoint(None) is None - assert PersistentResourceServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert PersistentResourceServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert PersistentResourceServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PersistentResourceServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - PersistentResourceServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert PersistentResourceServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - PersistentResourceServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert PersistentResourceServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - PersistentResourceServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert PersistentResourceServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -296,40 +164,6 @@ def test__get_client_cert_source(): assert PersistentResourceServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert PersistentResourceServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(PersistentResourceServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PersistentResourceServiceClient)) -@mock.patch.object(PersistentResourceServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PersistentResourceServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = PersistentResourceServiceClient._DEFAULT_UNIVERSE - default_endpoint = PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = PersistentResourceServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert PersistentResourceServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert PersistentResourceServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - assert PersistentResourceServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert PersistentResourceServiceClient._get_api_endpoint(None, None, default_universe, "always") == PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - assert PersistentResourceServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == PersistentResourceServiceClient.DEFAULT_MTLS_ENDPOINT - assert PersistentResourceServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert PersistentResourceServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - PersistentResourceServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert PersistentResourceServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert PersistentResourceServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert PersistentResourceServiceClient._get_universe_domain(None, None) == PersistentResourceServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - PersistentResourceServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -714,11 +548,12 @@ def test_persistent_resource_service_client_get_mtls_endpoint_and_cert_source(cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -761,11 +596,12 @@ def test_persistent_resource_service_client_get_mtls_endpoint_and_cert_source(cl for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3193,7 +3029,12 @@ def test_create_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped assert "persistentResourceId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreatePersistentResource, + "_BaseCreatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -3203,10 +3044,8 @@ def test_create_persistent_resource_rest_required_fields(request_type=persistent jsonified_request["parent"] = 'parent_value' jsonified_request["persistentResourceId"] = 'persistent_resource_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_persistent_resource._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("persistent_resource_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("persistentResourceId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3260,13 +3099,6 @@ def test_create_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_create_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(("persistentResourceId", )) & set(("parent", "persistentResource", "persistentResourceId", ))) - - def test_create_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3371,16 +3203,18 @@ def test_get_persistent_resource_rest_required_fields(request_type=persistent_re # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetPersistentResource, + "_BaseGetPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_persistent_resource._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3429,13 +3263,6 @@ def test_get_persistent_resource_rest_required_fields(request_type=persistent_re assert sorted(expected_params) == sorted(actual_params) -def test_get_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3538,17 +3365,20 @@ def test_list_persistent_resources_rest_required_fields(request_type=persistent_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_persistent_resources._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListPersistentResources, + "_BaseListPersistentResources__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_persistent_resources._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3598,13 +3428,6 @@ def test_list_persistent_resources_rest_required_fields(request_type=persistent_ assert sorted(expected_params) == sorted(actual_params) -def test_list_persistent_resources_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_persistent_resources._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_persistent_resources_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3776,16 +3599,18 @@ def test_delete_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeletePersistentResource, + "_BaseDeletePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_persistent_resource._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3831,13 +3656,6 @@ def test_delete_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_delete_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3941,15 +3759,18 @@ def test_update_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdatePersistentResource, + "_BaseUpdatePersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_persistent_resource._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3995,13 +3816,6 @@ def test_update_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_update_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("persistentResource", "updateMask", ))) - - def test_update_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4108,16 +3922,18 @@ def test_reboot_persistent_resource_rest_required_fields(request_type=persistent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).reboot_persistent_resource._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRebootPersistentResource, + "_BaseRebootPersistentResource__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).reboot_persistent_resource._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4164,13 +3980,6 @@ def test_reboot_persistent_resource_rest_required_fields(request_type=persistent assert sorted(expected_params) == sorted(actual_params) -def test_reboot_persistent_resource_rest_unset_required_fields(): - transport = transports.PersistentResourceServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.reboot_persistent_resource._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_reboot_persistent_resource_rest_flattened(): client = PersistentResourceServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_pipeline_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_pipeline_service.py index 4a099753ad..273fee6a68 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_pipeline_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_pipeline_service.py @@ -110,6 +110,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -153,150 +165,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert PipelineServiceClient._get_default_mtls_endpoint(None) is None - assert PipelineServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert PipelineServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert PipelineServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PipelineServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - PipelineServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert PipelineServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert PipelineServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert PipelineServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - PipelineServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert PipelineServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert PipelineServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - PipelineServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert PipelineServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert PipelineServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -310,40 +178,6 @@ def test__get_client_cert_source(): assert PipelineServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert PipelineServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(PipelineServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PipelineServiceClient)) -@mock.patch.object(PipelineServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PipelineServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = PipelineServiceClient._DEFAULT_UNIVERSE - default_endpoint = PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = PipelineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert PipelineServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert PipelineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - assert PipelineServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert PipelineServiceClient._get_api_endpoint(None, None, default_universe, "always") == PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - assert PipelineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == PipelineServiceClient.DEFAULT_MTLS_ENDPOINT - assert PipelineServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert PipelineServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - PipelineServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert PipelineServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert PipelineServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert PipelineServiceClient._get_universe_domain(None, None) == PipelineServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - PipelineServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -728,11 +562,12 @@ def test_pipeline_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -775,11 +610,12 @@ def test_pipeline_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5376,16 +5212,18 @@ def test_create_training_pipeline_rest_required_fields(request_type=pipeline_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTrainingPipeline, + "_BaseCreateTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -5435,13 +5273,6 @@ def test_create_training_pipeline_rest_required_fields(request_type=pipeline_ser assert sorted(expected_params) == sorted(actual_params) -def test_create_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "trainingPipeline", ))) - - def test_create_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5546,16 +5377,18 @@ def test_get_training_pipeline_rest_required_fields(request_type=pipeline_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTrainingPipeline, + "_BaseGetTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -5604,13 +5437,6 @@ def test_get_training_pipeline_rest_required_fields(request_type=pipeline_servic assert sorted(expected_params) == sorted(actual_params) -def test_get_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5713,17 +5539,20 @@ def test_list_training_pipelines_rest_required_fields(request_type=pipeline_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_training_pipelines._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTrainingPipelines, + "_BaseListTrainingPipelines__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_training_pipelines._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -5773,13 +5602,6 @@ def test_list_training_pipelines_rest_required_fields(request_type=pipeline_serv assert sorted(expected_params) == sorted(actual_params) -def test_list_training_pipelines_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_training_pipelines._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_training_pipelines_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5951,16 +5773,18 @@ def test_delete_training_pipeline_rest_required_fields(request_type=pipeline_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTrainingPipeline, + "_BaseDeleteTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6006,13 +5830,6 @@ def test_delete_training_pipeline_rest_required_fields(request_type=pipeline_ser assert sorted(expected_params) == sorted(actual_params) -def test_delete_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6113,16 +5930,18 @@ def test_cancel_training_pipeline_rest_required_fields(request_type=pipeline_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_training_pipeline._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelTrainingPipeline, + "_BaseCancelTrainingPipeline__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_training_pipeline._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6169,13 +5988,6 @@ def test_cancel_training_pipeline_rest_required_fields(request_type=pipeline_ser assert sorted(expected_params) == sorted(actual_params) -def test_cancel_training_pipeline_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_training_pipeline._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_training_pipeline_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6276,17 +6088,20 @@ def test_create_pipeline_job_rest_required_fields(request_type=pipeline_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreatePipelineJob, + "_BaseCreatePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_pipeline_job._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("pipeline_job_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pipelineJobId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6337,13 +6152,6 @@ def test_create_pipeline_job_rest_required_fields(request_type=pipeline_service. assert sorted(expected_params) == sorted(actual_params) -def test_create_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pipelineJobId", )) & set(("parent", "pipelineJob", ))) - - def test_create_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6450,16 +6258,18 @@ def test_get_pipeline_job_rest_required_fields(request_type=pipeline_service.Get # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetPipelineJob, + "_BaseGetPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_pipeline_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6508,13 +6318,6 @@ def test_get_pipeline_job_rest_required_fields(request_type=pipeline_service.Get assert sorted(expected_params) == sorted(actual_params) -def test_get_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6617,17 +6420,20 @@ def test_list_pipeline_jobs_rest_required_fields(request_type=pipeline_service.L # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_pipeline_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListPipelineJobs, + "_BaseListPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_pipeline_jobs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6677,13 +6483,6 @@ def test_list_pipeline_jobs_rest_required_fields(request_type=pipeline_service.L assert sorted(expected_params) == sorted(actual_params) -def test_list_pipeline_jobs_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_pipeline_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_pipeline_jobs_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6855,16 +6654,18 @@ def test_delete_pipeline_job_rest_required_fields(request_type=pipeline_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeletePipelineJob, + "_BaseDeletePipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_pipeline_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6910,13 +6711,6 @@ def test_delete_pipeline_job_rest_required_fields(request_type=pipeline_service. assert sorted(expected_params) == sorted(actual_params) -def test_delete_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7022,7 +6816,12 @@ def test_batch_delete_pipeline_jobs_rest_required_fields(request_type=pipeline_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_pipeline_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchDeletePipelineJobs, + "_BaseBatchDeletePipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -7030,9 +6829,6 @@ def test_batch_delete_pipeline_jobs_rest_required_fields(request_type=pipeline_s jsonified_request["parent"] = 'parent_value' jsonified_request["names"] = 'names_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_pipeline_jobs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -7081,13 +6877,6 @@ def test_batch_delete_pipeline_jobs_rest_required_fields(request_type=pipeline_s assert sorted(expected_params) == sorted(actual_params) -def test_batch_delete_pipeline_jobs_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_delete_pipeline_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "names", ))) - - def test_batch_delete_pipeline_jobs_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7190,16 +6979,18 @@ def test_cancel_pipeline_job_rest_required_fields(request_type=pipeline_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_pipeline_job._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelPipelineJob, + "_BaseCancelPipelineJob__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_pipeline_job._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7246,13 +7037,6 @@ def test_cancel_pipeline_job_rest_required_fields(request_type=pipeline_service. assert sorted(expected_params) == sorted(actual_params) -def test_cancel_pipeline_job_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_pipeline_job._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_pipeline_job_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7358,7 +7142,12 @@ def test_batch_cancel_pipeline_jobs_rest_required_fields(request_type=pipeline_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_cancel_pipeline_jobs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCancelPipelineJobs, + "_BaseBatchCancelPipelineJobs__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -7366,9 +7155,6 @@ def test_batch_cancel_pipeline_jobs_rest_required_fields(request_type=pipeline_s jsonified_request["parent"] = 'parent_value' jsonified_request["names"] = 'names_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_cancel_pipeline_jobs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -7417,13 +7203,6 @@ def test_batch_cancel_pipeline_jobs_rest_required_fields(request_type=pipeline_s assert sorted(expected_params) == sorted(actual_params) -def test_batch_cancel_pipeline_jobs_rest_unset_required_fields(): - transport = transports.PipelineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_cancel_pipeline_jobs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "names", ))) - - def test_batch_cancel_pipeline_jobs_rest_flattened(): client = PipelineServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_prediction_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_prediction_service.py index 07f47e5a9d..99287d8666 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_prediction_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_prediction_service.py @@ -92,6 +92,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -135,150 +147,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert PredictionServiceClient._get_default_mtls_endpoint(None) is None - assert PredictionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert PredictionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert PredictionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PredictionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - PredictionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert PredictionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert PredictionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert PredictionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - PredictionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert PredictionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert PredictionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - PredictionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert PredictionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert PredictionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -292,40 +160,6 @@ def test__get_client_cert_source(): assert PredictionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert PredictionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(PredictionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PredictionServiceClient)) -@mock.patch.object(PredictionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(PredictionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = PredictionServiceClient._DEFAULT_UNIVERSE - default_endpoint = PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = PredictionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert PredictionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert PredictionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - assert PredictionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert PredictionServiceClient._get_api_endpoint(None, None, default_universe, "always") == PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - assert PredictionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == PredictionServiceClient.DEFAULT_MTLS_ENDPOINT - assert PredictionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert PredictionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - PredictionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert PredictionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert PredictionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert PredictionServiceClient._get_universe_domain(None, None) == PredictionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - PredictionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -710,11 +544,12 @@ def test_prediction_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -757,11 +592,12 @@ def test_prediction_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5033,16 +4869,18 @@ def test_predict_rest_required_fields(request_type=prediction_service.PredictReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePredict, + "_BasePredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5092,13 +4930,6 @@ def test_predict_rest_required_fields(request_type=prediction_service.PredictReq assert sorted(expected_params) == sorted(actual_params) -def test_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "instances", ))) - - def test_predict_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5205,16 +5036,18 @@ def test_raw_predict_rest_required_fields(request_type=prediction_service.RawPre # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).raw_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRawPredict, + "_BaseRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).raw_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5262,13 +5095,6 @@ def test_raw_predict_rest_required_fields(request_type=prediction_service.RawPre assert sorted(expected_params) == sorted(actual_params) -def test_raw_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.raw_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_raw_predict_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5371,16 +5197,18 @@ def test_stream_raw_predict_rest_required_fields(request_type=prediction_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_raw_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamRawPredict, + "_BaseStreamRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_raw_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5431,13 +5259,6 @@ def test_stream_raw_predict_rest_required_fields(request_type=prediction_service assert sorted(expected_params) == sorted(actual_params) -def test_stream_raw_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stream_raw_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_stream_raw_predict_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5543,16 +5364,18 @@ def test_direct_predict_rest_required_fields(request_type=prediction_service.Dir # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDirectPredict, + "_BaseDirectPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5602,13 +5425,6 @@ def test_direct_predict_rest_required_fields(request_type=prediction_service.Dir assert sorted(expected_params) == sorted(actual_params) -def test_direct_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.direct_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_direct_raw_predict_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -5657,16 +5473,18 @@ def test_direct_raw_predict_rest_required_fields(request_type=prediction_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_raw_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDirectRawPredict, + "_BaseDirectRawPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).direct_raw_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5716,13 +5534,6 @@ def test_direct_raw_predict_rest_required_fields(request_type=prediction_service assert sorted(expected_params) == sorted(actual_params) -def test_direct_raw_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.direct_raw_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_stream_direct_predict_rest_no_http_options(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5804,16 +5615,18 @@ def test_server_streaming_predict_rest_required_fields(request_type=prediction_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).server_streaming_predict._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseServerStreamingPredict, + "_BaseServerStreamingPredict__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).server_streaming_predict._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5866,13 +5679,6 @@ def test_server_streaming_predict_rest_required_fields(request_type=prediction_s assert sorted(expected_params) == sorted(actual_params) -def test_server_streaming_predict_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.server_streaming_predict._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_streaming_raw_predict_rest_no_http_options(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5932,16 +5738,18 @@ def test_explain_rest_required_fields(request_type=prediction_service.ExplainReq # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).explain._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExplain, + "_BaseExplain__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).explain._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -5991,13 +5799,6 @@ def test_explain_rest_required_fields(request_type=prediction_service.ExplainReq assert sorted(expected_params) == sorted(actual_params) -def test_explain_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.explain._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", "instances", ))) - - def test_explain_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6106,16 +5907,18 @@ def test_count_tokens_rest_required_fields(request_type=prediction_service.Count # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).count_tokens._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCountTokens, + "_BaseCountTokens__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).count_tokens._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -6165,13 +5968,6 @@ def test_count_tokens_rest_required_fields(request_type=prediction_service.Count assert sorted(expected_params) == sorted(actual_params) -def test_count_tokens_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.count_tokens._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_count_tokens_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6276,16 +6072,18 @@ def test_generate_content_rest_required_fields(request_type=prediction_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGenerateContent, + "_BaseGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["model"] = 'model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).generate_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "model" in jsonified_request assert jsonified_request["model"] == 'model_value' @@ -6335,13 +6133,6 @@ def test_generate_content_rest_required_fields(request_type=prediction_service.G assert sorted(expected_params) == sorted(actual_params) -def test_generate_content_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.generate_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("model", "contents", ))) - - def test_generate_content_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6446,16 +6237,18 @@ def test_stream_generate_content_rest_required_fields(request_type=prediction_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_generate_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamGenerateContent, + "_BaseStreamGenerateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["model"] = 'model_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_generate_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "model" in jsonified_request assert jsonified_request["model"] == 'model_value' @@ -6508,13 +6301,6 @@ def test_stream_generate_content_rest_required_fields(request_type=prediction_se assert sorted(expected_params) == sorted(actual_params) -def test_stream_generate_content_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stream_generate_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("model", "contents", ))) - - def test_stream_generate_content_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6622,16 +6408,18 @@ def test_chat_completions_rest_required_fields(request_type=prediction_service.C # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).chat_completions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseChatCompletions, + "_BaseChatCompletions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["endpoint"] = 'endpoint_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).chat_completions._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "endpoint" in jsonified_request assert jsonified_request["endpoint"] == 'endpoint_value' @@ -6682,13 +6470,6 @@ def test_chat_completions_rest_required_fields(request_type=prediction_service.C assert sorted(expected_params) == sorted(actual_params) -def test_chat_completions_rest_unset_required_fields(): - transport = transports.PredictionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.chat_completions._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("endpoint", ))) - - def test_chat_completions_rest_flattened(): client = PredictionServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_execution_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_execution_service.py index 702bd406bc..c3076defb8 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_execution_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_execution_service.py @@ -85,6 +85,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -128,150 +140,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(None) is None - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ReasoningEngineExecutionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ReasoningEngineExecutionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineExecutionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ReasoningEngineExecutionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ReasoningEngineExecutionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ReasoningEngineExecutionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -285,40 +153,6 @@ def test__get_client_cert_source(): assert ReasoningEngineExecutionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ReasoningEngineExecutionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ReasoningEngineExecutionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineExecutionServiceClient)) -@mock.patch.object(ReasoningEngineExecutionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineExecutionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - default_endpoint = ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ReasoningEngineExecutionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, default_universe, "always") == ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ReasoningEngineExecutionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ReasoningEngineExecutionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineExecutionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ReasoningEngineExecutionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ReasoningEngineExecutionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ReasoningEngineExecutionServiceClient._get_universe_domain(None, None) == ReasoningEngineExecutionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ReasoningEngineExecutionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -703,11 +537,12 @@ def test_reasoning_engine_execution_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -750,11 +585,12 @@ def test_reasoning_engine_execution_service_client_get_mtls_endpoint_and_cert_so for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -1949,16 +1785,18 @@ def test_query_reasoning_engine_rest_required_fields(request_type=reasoning_engi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseQueryReasoningEngine, + "_BaseQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2008,13 +1846,6 @@ def test_query_reasoning_engine_rest_required_fields(request_type=reasoning_engi assert sorted(expected_params) == sorted(actual_params) -def test_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_stream_query_reasoning_engine_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -2063,16 +1894,18 @@ def test_stream_query_reasoning_engine_rest_required_fields(request_type=reasoni # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStreamQueryReasoningEngine, + "_BaseStreamQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stream_query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2123,13 +1956,6 @@ def test_stream_query_reasoning_engine_rest_required_fields(request_type=reasoni assert sorted(expected_params) == sorted(actual_params) -def test_stream_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stream_query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_async_query_reasoning_engine_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -2182,16 +2008,18 @@ def test_async_query_reasoning_engine_rest_required_fields(request_type=reasonin # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAsyncQueryReasoningEngine, + "_BaseAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2238,13 +2066,6 @@ def test_async_query_reasoning_engine_rest_required_fields(request_type=reasonin assert sorted(expected_params) == sorted(actual_params) -def test_async_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.async_query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_cancel_async_query_reasoning_engine_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -2294,7 +2115,12 @@ def test_cancel_async_query_reasoning_engine_rest_required_fields(request_type=r # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_async_query_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCancelAsyncQueryReasoningEngine, + "_BaseCancelAsyncQueryReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -2302,9 +2128,6 @@ def test_cancel_async_query_reasoning_engine_rest_required_fields(request_type=r jsonified_request["name"] = 'name_value' jsonified_request["operationName"] = 'operation_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).cancel_async_query_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2356,13 +2179,6 @@ def test_cancel_async_query_reasoning_engine_rest_required_fields(request_type=r assert sorted(expected_params) == sorted(actual_params) -def test_cancel_async_query_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineExecutionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.cancel_async_query_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "operationName", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.ReasoningEngineExecutionServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_runtime_revision_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_runtime_revision_service.py index f6ff96f0eb..ad75880cd5 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_runtime_revision_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_runtime_revision_service.py @@ -87,6 +87,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -130,150 +142,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ReasoningEngineRuntimeRevisionServiceClient._get_default_mtls_endpoint(None) is None - assert ReasoningEngineRuntimeRevisionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ReasoningEngineRuntimeRevisionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ReasoningEngineRuntimeRevisionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineRuntimeRevisionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineRuntimeRevisionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ReasoningEngineRuntimeRevisionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ReasoningEngineRuntimeRevisionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ReasoningEngineRuntimeRevisionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -287,40 +155,6 @@ def test__get_client_cert_source(): assert ReasoningEngineRuntimeRevisionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ReasoningEngineRuntimeRevisionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ReasoningEngineRuntimeRevisionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineRuntimeRevisionServiceClient)) -@mock.patch.object(ReasoningEngineRuntimeRevisionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineRuntimeRevisionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_UNIVERSE - default_endpoint = ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ReasoningEngineRuntimeRevisionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(None, None, default_universe, "always") == ReasoningEngineRuntimeRevisionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ReasoningEngineRuntimeRevisionServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineRuntimeRevisionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ReasoningEngineRuntimeRevisionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ReasoningEngineRuntimeRevisionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ReasoningEngineRuntimeRevisionServiceClient._get_universe_domain(None, None) == ReasoningEngineRuntimeRevisionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ReasoningEngineRuntimeRevisionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -705,11 +539,12 @@ def test_reasoning_engine_runtime_revision_service_client_get_mtls_endpoint_and_ for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -752,11 +587,12 @@ def test_reasoning_engine_runtime_revision_service_client_get_mtls_endpoint_and_ for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2179,16 +2015,18 @@ def test_get_reasoning_engine_runtime_revision_rest_required_fields(request_type # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_reasoning_engine_runtime_revision._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetReasoningEngineRuntimeRevision, + "_BaseGetReasoningEngineRuntimeRevision__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_reasoning_engine_runtime_revision._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2237,13 +2075,6 @@ def test_get_reasoning_engine_runtime_revision_rest_required_fields(request_type assert sorted(expected_params) == sorted(actual_params) -def test_get_reasoning_engine_runtime_revision_rest_unset_required_fields(): - transport = transports.ReasoningEngineRuntimeRevisionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_reasoning_engine_runtime_revision._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_reasoning_engine_runtime_revision_rest_flattened(): client = ReasoningEngineRuntimeRevisionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2346,17 +2177,20 @@ def test_list_reasoning_engine_runtime_revisions_rest_required_fields(request_ty # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_reasoning_engine_runtime_revisions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListReasoningEngineRuntimeRevisions, + "_BaseListReasoningEngineRuntimeRevisions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_reasoning_engine_runtime_revisions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -2406,13 +2240,6 @@ def test_list_reasoning_engine_runtime_revisions_rest_required_fields(request_ty assert sorted(expected_params) == sorted(actual_params) -def test_list_reasoning_engine_runtime_revisions_rest_unset_required_fields(): - transport = transports.ReasoningEngineRuntimeRevisionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_reasoning_engine_runtime_revisions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_reasoning_engine_runtime_revisions_rest_flattened(): client = ReasoningEngineRuntimeRevisionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2584,16 +2411,18 @@ def test_delete_reasoning_engine_runtime_revision_rest_required_fields(request_t # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_reasoning_engine_runtime_revision._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteReasoningEngineRuntimeRevision, + "_BaseDeleteReasoningEngineRuntimeRevision__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_reasoning_engine_runtime_revision._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -2639,13 +2468,6 @@ def test_delete_reasoning_engine_runtime_revision_rest_required_fields(request_t assert sorted(expected_params) == sorted(actual_params) -def test_delete_reasoning_engine_runtime_revision_rest_unset_required_fields(): - transport = transports.ReasoningEngineRuntimeRevisionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_reasoning_engine_runtime_revision._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_reasoning_engine_runtime_revision_rest_flattened(): client = ReasoningEngineRuntimeRevisionServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_service.py index 1e3c4bf65b..25440a4556 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_reasoning_engine_service.py @@ -94,6 +94,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -137,150 +149,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(None) is None - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ReasoningEngineServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ReasoningEngineServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ReasoningEngineServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ReasoningEngineServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ReasoningEngineServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ReasoningEngineServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -294,40 +162,6 @@ def test__get_client_cert_source(): assert ReasoningEngineServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ReasoningEngineServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ReasoningEngineServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineServiceClient)) -@mock.patch.object(ReasoningEngineServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ReasoningEngineServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ReasoningEngineServiceClient._DEFAULT_UNIVERSE - default_endpoint = ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ReasoningEngineServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ReasoningEngineServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ReasoningEngineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, default_universe, "always") == ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ReasoningEngineServiceClient.DEFAULT_MTLS_ENDPOINT - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ReasoningEngineServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ReasoningEngineServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ReasoningEngineServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ReasoningEngineServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ReasoningEngineServiceClient._get_universe_domain(None, None) == ReasoningEngineServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ReasoningEngineServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -712,11 +546,12 @@ def test_reasoning_engine_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -759,11 +594,12 @@ def test_reasoning_engine_service_client_get_mtls_endpoint_and_cert_source(clien for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2850,16 +2686,18 @@ def test_create_reasoning_engine_rest_required_fields(request_type=reasoning_eng # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateReasoningEngine, + "_BaseCreateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2906,13 +2744,6 @@ def test_create_reasoning_engine_rest_required_fields(request_type=reasoning_eng assert sorted(expected_params) == sorted(actual_params) -def test_create_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "reasoningEngine", ))) - - def test_create_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3015,16 +2846,18 @@ def test_get_reasoning_engine_rest_required_fields(request_type=reasoning_engine # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetReasoningEngine, + "_BaseGetReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_reasoning_engine._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3073,13 +2906,6 @@ def test_get_reasoning_engine_rest_required_fields(request_type=reasoning_engine assert sorted(expected_params) == sorted(actual_params) -def test_get_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3182,17 +3008,20 @@ def test_list_reasoning_engines_rest_required_fields(request_type=reasoning_engi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_reasoning_engines._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListReasoningEngines, + "_BaseListReasoningEngines__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_reasoning_engines._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3242,13 +3071,6 @@ def test_list_reasoning_engines_rest_required_fields(request_type=reasoning_engi assert sorted(expected_params) == sorted(actual_params) -def test_list_reasoning_engines_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_reasoning_engines._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_reasoning_engines_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3419,15 +3241,18 @@ def test_update_reasoning_engine_rest_required_fields(request_type=reasoning_eng # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateReasoningEngine, + "_BaseUpdateReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_reasoning_engine._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3473,13 +3298,6 @@ def test_update_reasoning_engine_rest_required_fields(request_type=reasoning_eng assert sorted(expected_params) == sorted(actual_params) -def test_update_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("reasoningEngine", ))) - - def test_update_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3586,17 +3404,20 @@ def test_delete_reasoning_engine_rest_required_fields(request_type=reasoning_eng # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_reasoning_engine._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteReasoningEngine, + "_BaseDeleteReasoningEngine__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_reasoning_engine._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -3643,13 +3464,6 @@ def test_delete_reasoning_engine_rest_required_fields(request_type=reasoning_eng assert sorted(expected_params) == sorted(actual_params) -def test_delete_reasoning_engine_rest_unset_required_fields(): - transport = transports.ReasoningEngineServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_reasoning_engine._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_reasoning_engine_rest_flattened(): client = ReasoningEngineServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_schedule_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_schedule_service.py index 6cb67dd10b..abc587448c 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_schedule_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_schedule_service.py @@ -119,6 +119,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -162,150 +174,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert ScheduleServiceClient._get_default_mtls_endpoint(None) is None - assert ScheduleServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert ScheduleServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert ScheduleServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ScheduleServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - ScheduleServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert ScheduleServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - ScheduleServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert ScheduleServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert ScheduleServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - ScheduleServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert ScheduleServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert ScheduleServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -319,40 +187,6 @@ def test__get_client_cert_source(): assert ScheduleServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert ScheduleServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(ScheduleServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ScheduleServiceClient)) -@mock.patch.object(ScheduleServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(ScheduleServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = ScheduleServiceClient._DEFAULT_UNIVERSE - default_endpoint = ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = ScheduleServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert ScheduleServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert ScheduleServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - assert ScheduleServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert ScheduleServiceClient._get_api_endpoint(None, None, default_universe, "always") == ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - assert ScheduleServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == ScheduleServiceClient.DEFAULT_MTLS_ENDPOINT - assert ScheduleServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert ScheduleServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - ScheduleServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert ScheduleServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert ScheduleServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert ScheduleServiceClient._get_universe_domain(None, None) == ScheduleServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - ScheduleServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -737,11 +571,12 @@ def test_schedule_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -784,11 +619,12 @@ def test_schedule_service_client_get_mtls_endpoint_and_cert_source(client_class) for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3560,16 +3396,18 @@ def test_create_schedule_rest_required_fields(request_type=schedule_service.Crea # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateSchedule, + "_BaseCreateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3619,13 +3457,6 @@ def test_create_schedule_rest_required_fields(request_type=schedule_service.Crea assert sorted(expected_params) == sorted(actual_params) -def test_create_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "schedule", ))) - - def test_create_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3734,16 +3565,18 @@ def test_delete_schedule_rest_required_fields(request_type=schedule_service.Dele # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSchedule, + "_BaseDeleteSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3789,13 +3622,6 @@ def test_delete_schedule_rest_required_fields(request_type=schedule_service.Dele assert sorted(expected_params) == sorted(actual_params) -def test_delete_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3896,16 +3722,18 @@ def test_get_schedule_rest_required_fields(request_type=schedule_service.GetSche # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetSchedule, + "_BaseGetSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3954,13 +3782,6 @@ def test_get_schedule_rest_required_fields(request_type=schedule_service.GetSche assert sorted(expected_params) == sorted(actual_params) -def test_get_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4063,17 +3884,20 @@ def test_list_schedules_rest_required_fields(request_type=schedule_service.ListS # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_schedules._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSchedules, + "_BaseListSchedules__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_schedules._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4123,13 +3947,6 @@ def test_list_schedules_rest_required_fields(request_type=schedule_service.ListS assert sorted(expected_params) == sorted(actual_params) -def test_list_schedules_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_schedules._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_schedules_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4297,16 +4114,18 @@ def test_pause_schedule_rest_required_fields(request_type=schedule_service.Pause # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BasePauseSchedule, + "_BasePauseSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).pause_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4353,13 +4172,6 @@ def test_pause_schedule_rest_required_fields(request_type=schedule_service.Pause assert sorted(expected_params) == sorted(actual_params) -def test_pause_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.pause_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_pause_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4460,16 +4272,18 @@ def test_resume_schedule_rest_required_fields(request_type=schedule_service.Resu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseResumeSchedule, + "_BaseResumeSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).resume_schedule._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4516,13 +4330,6 @@ def test_resume_schedule_rest_required_fields(request_type=schedule_service.Resu assert sorted(expected_params) == sorted(actual_params) -def test_resume_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.resume_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_resume_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4624,15 +4431,18 @@ def test_update_schedule_rest_required_fields(request_type=schedule_service.Upda # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_schedule._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateSchedule, + "_BaseUpdateSchedule__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_schedule._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4681,13 +4491,6 @@ def test_update_schedule_rest_required_fields(request_type=schedule_service.Upda assert sorted(expected_params) == sorted(actual_params) -def test_update_schedule_rest_unset_required_fields(): - transport = transports.ScheduleServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_schedule._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("schedule", "updateMask", ))) - - def test_update_schedule_rest_flattened(): client = ScheduleServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_session_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_session_service.py index 67a050dc37..1dc128f168 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_session_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_session_service.py @@ -94,6 +94,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -137,150 +149,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert SessionServiceClient._get_default_mtls_endpoint(None) is None - assert SessionServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert SessionServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert SessionServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert SessionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SessionServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SessionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - SessionServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert SessionServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert SessionServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert SessionServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert SessionServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - SessionServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert SessionServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert SessionServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - SessionServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert SessionServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert SessionServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -294,40 +162,6 @@ def test__get_client_cert_source(): assert SessionServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert SessionServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(SessionServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SessionServiceClient)) -@mock.patch.object(SessionServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SessionServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = SessionServiceClient._DEFAULT_UNIVERSE - default_endpoint = SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = SessionServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert SessionServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert SessionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == SessionServiceClient.DEFAULT_MTLS_ENDPOINT - assert SessionServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert SessionServiceClient._get_api_endpoint(None, None, default_universe, "always") == SessionServiceClient.DEFAULT_MTLS_ENDPOINT - assert SessionServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == SessionServiceClient.DEFAULT_MTLS_ENDPOINT - assert SessionServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert SessionServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - SessionServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert SessionServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert SessionServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert SessionServiceClient._get_universe_domain(None, None) == SessionServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - SessionServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -712,11 +546,12 @@ def test_session_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -759,11 +594,12 @@ def test_session_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -3689,17 +3525,20 @@ def test_create_session_rest_required_fields(request_type=session_service.Create # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateSession, + "_BaseCreateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_session._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("session_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("sessionId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3747,13 +3586,6 @@ def test_create_session_rest_required_fields(request_type=session_service.Create assert sorted(expected_params) == sorted(actual_params) -def test_create_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(("sessionId", )) & set(("parent", "session", ))) - - def test_create_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3858,16 +3690,18 @@ def test_get_session_rest_required_fields(request_type=session_service.GetSessio # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetSession, + "_BaseGetSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_session._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3916,13 +3750,6 @@ def test_get_session_rest_required_fields(request_type=session_service.GetSessio assert sorted(expected_params) == sorted(actual_params) -def test_get_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4025,17 +3852,20 @@ def test_list_sessions_rest_required_fields(request_type=session_service.ListSes # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_sessions._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSessions, + "_BaseListSessions__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_sessions._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4085,13 +3915,6 @@ def test_list_sessions_rest_required_fields(request_type=session_service.ListSes assert sorted(expected_params) == sorted(actual_params) -def test_list_sessions_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_sessions._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_sessions_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4258,15 +4081,18 @@ def test_update_session_rest_required_fields(request_type=session_service.Update # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateSession, + "_BaseUpdateSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_session._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -4315,13 +4141,6 @@ def test_update_session_rest_required_fields(request_type=session_service.Update assert sorted(expected_params) == sorted(actual_params) -def test_update_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("session", ))) - - def test_update_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4430,16 +4249,18 @@ def test_delete_session_rest_required_fields(request_type=session_service.Delete # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_session._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSession, + "_BaseDeleteSession__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_session._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4485,13 +4306,6 @@ def test_delete_session_rest_required_fields(request_type=session_service.Delete assert sorted(expected_params) == sorted(actual_params) -def test_delete_session_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_session._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_session_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4592,17 +4406,20 @@ def test_list_events_rest_required_fields(request_type=session_service.ListEvent # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_events._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListEvents, + "_BaseListEvents__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_events._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -4652,13 +4469,6 @@ def test_list_events_rest_required_fields(request_type=session_service.ListEvent assert sorted(expected_params) == sorted(actual_params) -def test_list_events_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_events._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_events_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -4826,16 +4636,18 @@ def test_append_event_rest_required_fields(request_type=session_service.AppendEv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).append_event._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAppendEvent, + "_BaseAppendEvent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).append_event._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -4885,13 +4697,6 @@ def test_append_event_rest_required_fields(request_type=session_service.AppendEv assert sorted(expected_params) == sorted(actual_params) -def test_append_event_rest_unset_required_fields(): - transport = transports.SessionServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.append_event._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", "event", ))) - - def test_append_event_rest_flattened(): client = SessionServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_specialist_pool_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_specialist_pool_service.py index d8e026ee6f..6ba2d6c010 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_specialist_pool_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_specialist_pool_service.py @@ -88,6 +88,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -131,150 +143,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(None) is None - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert SpecialistPoolServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - SpecialistPoolServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert SpecialistPoolServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - SpecialistPoolServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert SpecialistPoolServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - SpecialistPoolServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert SpecialistPoolServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -288,40 +156,6 @@ def test__get_client_cert_source(): assert SpecialistPoolServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert SpecialistPoolServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(SpecialistPoolServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SpecialistPoolServiceClient)) -@mock.patch.object(SpecialistPoolServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(SpecialistPoolServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = SpecialistPoolServiceClient._DEFAULT_UNIVERSE - default_endpoint = SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = SpecialistPoolServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert SpecialistPoolServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert SpecialistPoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, default_universe, "always") == SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert SpecialistPoolServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == SpecialistPoolServiceClient.DEFAULT_MTLS_ENDPOINT - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert SpecialistPoolServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - SpecialistPoolServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert SpecialistPoolServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert SpecialistPoolServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert SpecialistPoolServiceClient._get_universe_domain(None, None) == SpecialistPoolServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - SpecialistPoolServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -706,11 +540,12 @@ def test_specialist_pool_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -753,11 +588,12 @@ def test_specialist_pool_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2850,16 +2686,18 @@ def test_create_specialist_pool_rest_required_fields(request_type=specialist_poo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateSpecialistPool, + "_BaseCreateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_specialist_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2906,13 +2744,6 @@ def test_create_specialist_pool_rest_required_fields(request_type=specialist_poo assert sorted(expected_params) == sorted(actual_params) -def test_create_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "specialistPool", ))) - - def test_create_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3015,16 +2846,18 @@ def test_get_specialist_pool_rest_required_fields(request_type=specialist_pool_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetSpecialistPool, + "_BaseGetSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_specialist_pool._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -3073,13 +2906,6 @@ def test_get_specialist_pool_rest_required_fields(request_type=specialist_pool_s assert sorted(expected_params) == sorted(actual_params) -def test_get_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3182,17 +3008,20 @@ def test_list_specialist_pools_rest_required_fields(request_type=specialist_pool # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_specialist_pools._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListSpecialistPools, + "_BaseListSpecialistPools__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_specialist_pools._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -3242,13 +3071,6 @@ def test_list_specialist_pools_rest_required_fields(request_type=specialist_pool assert sorted(expected_params) == sorted(actual_params) -def test_list_specialist_pools_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_specialist_pools._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_specialist_pools_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3420,17 +3242,20 @@ def test_delete_specialist_pool_rest_required_fields(request_type=specialist_poo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteSpecialistPool, + "_BaseDeleteSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_specialist_pool._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. assert not set(unset_fields) - set(("force", )) - jsonified_request.update(unset_fields) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -3477,13 +3302,6 @@ def test_delete_specialist_pool_rest_required_fields(request_type=specialist_poo assert sorted(expected_params) == sorted(actual_params) -def test_delete_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", )) & set(("name", ))) - - def test_delete_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3587,15 +3405,18 @@ def test_update_specialist_pool_rest_required_fields(request_type=specialist_poo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_specialist_pool._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateSpecialistPool, + "_BaseUpdateSpecialistPool__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_specialist_pool._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -3641,13 +3462,6 @@ def test_update_specialist_pool_rest_required_fields(request_type=specialist_poo assert sorted(expected_params) == sorted(actual_params) -def test_update_specialist_pool_rest_unset_required_fields(): - transport = transports.SpecialistPoolServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_specialist_pool._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("specialistPool", "updateMask", ))) - - def test_update_specialist_pool_rest_flattened(): client = SpecialistPoolServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_tensorboard_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_tensorboard_service.py index 8079ab0cfa..c50331382f 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_tensorboard_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_tensorboard_service.py @@ -97,6 +97,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -140,150 +152,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert TensorboardServiceClient._get_default_mtls_endpoint(None) is None - assert TensorboardServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert TensorboardServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert TensorboardServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert TensorboardServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - TensorboardServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert TensorboardServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - TensorboardServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert TensorboardServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert TensorboardServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - TensorboardServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert TensorboardServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert TensorboardServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -297,40 +165,6 @@ def test__get_client_cert_source(): assert TensorboardServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert TensorboardServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(TensorboardServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(TensorboardServiceClient)) -@mock.patch.object(TensorboardServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(TensorboardServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = TensorboardServiceClient._DEFAULT_UNIVERSE - default_endpoint = TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = TensorboardServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert TensorboardServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert TensorboardServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - assert TensorboardServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert TensorboardServiceClient._get_api_endpoint(None, None, default_universe, "always") == TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - assert TensorboardServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == TensorboardServiceClient.DEFAULT_MTLS_ENDPOINT - assert TensorboardServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert TensorboardServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - TensorboardServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert TensorboardServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert TensorboardServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert TensorboardServiceClient._get_universe_domain(None, None) == TensorboardServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - TensorboardServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -715,11 +549,12 @@ def test_tensorboard_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -762,11 +597,12 @@ def test_tensorboard_service_client_get_mtls_endpoint_and_cert_source(client_cla for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -11693,16 +11529,18 @@ def test_create_tensorboard_rest_required_fields(request_type=tensorboard_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboard, + "_BaseCreateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -11749,13 +11587,6 @@ def test_create_tensorboard_rest_required_fields(request_type=tensorboard_servic assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "tensorboard", ))) - - def test_create_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11858,16 +11689,18 @@ def test_get_tensorboard_rest_required_fields(request_type=tensorboard_service.G # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboard, + "_BaseGetTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11916,13 +11749,6 @@ def test_get_tensorboard_rest_required_fields(request_type=tensorboard_service.G assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12028,15 +11854,18 @@ def test_update_tensorboard_rest_required_fields(request_type=tensorboard_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboard, + "_BaseUpdateTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -12082,13 +11911,6 @@ def test_update_tensorboard_rest_required_fields(request_type=tensorboard_servic assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboard", ))) - - def test_update_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12191,17 +12013,20 @@ def test_list_tensorboards_rest_required_fields(request_type=tensorboard_service # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboards._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboards, + "_BaseListTensorboards__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboards._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12251,13 +12076,6 @@ def test_list_tensorboards_rest_required_fields(request_type=tensorboard_service assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboards_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboards._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboards_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12429,16 +12247,18 @@ def test_delete_tensorboard_rest_required_fields(request_type=tensorboard_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboard, + "_BaseDeleteTensorboard__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -12484,13 +12304,6 @@ def test_delete_tensorboard_rest_required_fields(request_type=tensorboard_servic assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12591,16 +12404,18 @@ def test_read_tensorboard_usage_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_usage._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardUsage, + "_BaseReadTensorboardUsage__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboard"] = 'tensorboard_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_usage._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboard" in jsonified_request assert jsonified_request["tensorboard"] == 'tensorboard_value' @@ -12649,13 +12464,6 @@ def test_read_tensorboard_usage_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_usage_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_usage._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboard", ))) - - def test_read_tensorboard_usage_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12758,16 +12566,18 @@ def test_read_tensorboard_size_rest_required_fields(request_type=tensorboard_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_size._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardSize, + "_BaseReadTensorboardSize__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboard"] = 'tensorboard_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_size._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboard" in jsonified_request assert jsonified_request["tensorboard"] == 'tensorboard_value' @@ -12816,13 +12626,6 @@ def test_read_tensorboard_size_rest_required_fields(request_type=tensorboard_ser assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_size_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_size._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboard", ))) - - def test_read_tensorboard_size_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12927,7 +12730,12 @@ def test_create_tensorboard_experiment_rest_required_fields(request_type=tensorb # verify fields with default values are dropped assert "tensorboardExperimentId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboardExperiment, + "_BaseCreateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -12937,10 +12745,8 @@ def test_create_tensorboard_experiment_rest_required_fields(request_type=tensorb jsonified_request["parent"] = 'parent_value' jsonified_request["tensorboardExperimentId"] = 'tensorboard_experiment_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_experiment._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("tensorboard_experiment_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("tensorboardExperimentId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12997,13 +12803,6 @@ def test_create_tensorboard_experiment_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(("tensorboardExperimentId", )) & set(("parent", "tensorboardExperimentId", ))) - - def test_create_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13110,16 +12909,18 @@ def test_get_tensorboard_experiment_rest_required_fields(request_type=tensorboar # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboardExperiment, + "_BaseGetTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_experiment._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13168,13 +12969,6 @@ def test_get_tensorboard_experiment_rest_required_fields(request_type=tensorboar assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13276,15 +13070,18 @@ def test_update_tensorboard_experiment_rest_required_fields(request_type=tensorb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboardExperiment, + "_BaseUpdateTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_experiment._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -13333,13 +13130,6 @@ def test_update_tensorboard_experiment_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboardExperiment", ))) - - def test_update_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13444,17 +13234,20 @@ def test_list_tensorboard_experiments_rest_required_fields(request_type=tensorbo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_experiments._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboardExperiments, + "_BaseListTensorboardExperiments__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_experiments._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13504,13 +13297,6 @@ def test_list_tensorboard_experiments_rest_required_fields(request_type=tensorbo assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboard_experiments_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboard_experiments._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboard_experiments_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13682,16 +13468,18 @@ def test_delete_tensorboard_experiment_rest_required_fields(request_type=tensorb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_experiment._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboardExperiment, + "_BaseDeleteTensorboardExperiment__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_experiment._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13737,13 +13525,6 @@ def test_delete_tensorboard_experiment_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_experiment_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard_experiment._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_experiment_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13846,7 +13627,12 @@ def test_create_tensorboard_run_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped assert "tensorboardRunId" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboardRun, + "_BaseCreateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -13856,10 +13642,8 @@ def test_create_tensorboard_run_rest_required_fields(request_type=tensorboard_se jsonified_request["parent"] = 'parent_value' jsonified_request["tensorboardRunId"] = 'tensorboard_run_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_run._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("tensorboard_run_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("tensorboardRunId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13916,13 +13700,6 @@ def test_create_tensorboard_run_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(("tensorboardRunId", )) & set(("parent", "tensorboardRun", "tensorboardRunId", ))) - - def test_create_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14029,16 +13806,18 @@ def test_batch_create_tensorboard_runs_rest_required_fields(request_type=tensorb # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_runs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateTensorboardRuns, + "_BaseBatchCreateTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_runs._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14088,13 +13867,6 @@ def test_batch_create_tensorboard_runs_rest_required_fields(request_type=tensorb assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_tensorboard_runs_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_tensorboard_runs._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_tensorboard_runs_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14199,16 +13971,18 @@ def test_get_tensorboard_run_rest_required_fields(request_type=tensorboard_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboardRun, + "_BaseGetTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_run._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14257,13 +14031,6 @@ def test_get_tensorboard_run_rest_required_fields(request_type=tensorboard_servi assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14365,15 +14132,18 @@ def test_update_tensorboard_run_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboardRun, + "_BaseUpdateTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_run._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -14422,13 +14192,6 @@ def test_update_tensorboard_run_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboardRun", ))) - - def test_update_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14533,17 +14296,20 @@ def test_list_tensorboard_runs_rest_required_fields(request_type=tensorboard_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_runs._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboardRuns, + "_BaseListTensorboardRuns__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_runs._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -14593,13 +14359,6 @@ def test_list_tensorboard_runs_rest_required_fields(request_type=tensorboard_ser assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboard_runs_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboard_runs._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboard_runs_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14771,16 +14530,18 @@ def test_delete_tensorboard_run_rest_required_fields(request_type=tensorboard_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_run._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboardRun, + "_BaseDeleteTensorboardRun__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_run._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -14826,13 +14587,6 @@ def test_delete_tensorboard_run_rest_required_fields(request_type=tensorboard_se assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_run_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard_run._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_run_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -14933,16 +14687,18 @@ def test_batch_create_tensorboard_time_series_rest_required_fields(request_type= # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateTensorboardTimeSeries, + "_BaseBatchCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_tensorboard_time_series._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -14992,13 +14748,6 @@ def test_batch_create_tensorboard_time_series_rest_required_fields(request_type= assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_batch_create_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15103,17 +14852,20 @@ def test_create_tensorboard_time_series_rest_required_fields(request_type=tensor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTensorboardTimeSeries, + "_BaseCreateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_tensorboard_time_series._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("tensorboard_time_series_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("tensorboardTimeSeriesId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -15164,13 +14916,6 @@ def test_create_tensorboard_time_series_rest_required_fields(request_type=tensor assert sorted(expected_params) == sorted(actual_params) -def test_create_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(("tensorboardTimeSeriesId", )) & set(("parent", "tensorboardTimeSeries", ))) - - def test_create_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15275,16 +15020,18 @@ def test_get_tensorboard_time_series_rest_required_fields(request_type=tensorboa # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTensorboardTimeSeries, + "_BaseGetTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_tensorboard_time_series._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15333,13 +15080,6 @@ def test_get_tensorboard_time_series_rest_required_fields(request_type=tensorboa assert sorted(expected_params) == sorted(actual_params) -def test_get_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15441,15 +15181,18 @@ def test_update_tensorboard_time_series_rest_required_fields(request_type=tensor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateTensorboardTimeSeries, + "_BaseUpdateTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_tensorboard_time_series._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("update_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("updateMask", )) # verify required fields with non-default values are left alone @@ -15498,13 +15241,6 @@ def test_update_tensorboard_time_series_rest_required_fields(request_type=tensor assert sorted(expected_params) == sorted(actual_params) -def test_update_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(("updateMask", )) & set(("updateMask", "tensorboardTimeSeries", ))) - - def test_update_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15609,17 +15345,20 @@ def test_list_tensorboard_time_series_rest_required_fields(request_type=tensorbo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTensorboardTimeSeries, + "_BaseListTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_tensorboard_time_series._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "order_by", "page_size", "page_token", "read_mask", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -15669,13 +15408,6 @@ def test_list_tensorboard_time_series_rest_required_fields(request_type=tensorbo assert sorted(expected_params) == sorted(actual_params) -def test_list_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "orderBy", "pageSize", "pageToken", "readMask", )) & set(("parent", ))) - - def test_list_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -15847,16 +15579,18 @@ def test_delete_tensorboard_time_series_rest_required_fields(request_type=tensor # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_time_series._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTensorboardTimeSeries, + "_BaseDeleteTensorboardTimeSeries__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_tensorboard_time_series._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -15902,13 +15636,6 @@ def test_delete_tensorboard_time_series_rest_required_fields(request_type=tensor assert sorted(expected_params) == sorted(actual_params) -def test_delete_tensorboard_time_series_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_tensorboard_time_series._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_tensorboard_time_series_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16011,7 +15738,12 @@ def test_batch_read_tensorboard_time_series_data_rest_required_fields(request_ty # verify fields with default values are dropped assert "timeSeries" not in jsonified_request - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchReadTensorboardTimeSeriesData, + "_BaseBatchReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -16021,10 +15753,8 @@ def test_batch_read_tensorboard_time_series_data_rest_required_fields(request_ty jsonified_request["tensorboard"] = 'tensorboard_value' jsonified_request["timeSeries"] = 'time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("time_series", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("timeSeries", )) # verify required fields with non-default values are left alone assert "tensorboard" in jsonified_request @@ -16080,13 +15810,6 @@ def test_batch_read_tensorboard_time_series_data_rest_required_fields(request_ty assert sorted(expected_params) == sorted(actual_params) -def test_batch_read_tensorboard_time_series_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_read_tensorboard_time_series_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(("timeSeries", )) & set(("tensorboard", "timeSeries", ))) - - def test_batch_read_tensorboard_time_series_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16189,17 +15912,20 @@ def test_read_tensorboard_time_series_data_rest_required_fields(request_type=ten # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardTimeSeriesData, + "_BaseReadTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardTimeSeries"] = 'tensorboard_time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("filter", "max_data_points", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("filter", "maxDataPoints", )) # verify required fields with non-default values are left alone assert "tensorboardTimeSeries" in jsonified_request @@ -16249,13 +15975,6 @@ def test_read_tensorboard_time_series_data_rest_required_fields(request_type=ten assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_time_series_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_time_series_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(("filter", "maxDataPoints", )) & set(("tensorboardTimeSeries", ))) - - def test_read_tensorboard_time_series_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16358,17 +16077,20 @@ def test_read_tensorboard_blob_data_rest_required_fields(request_type=tensorboar # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_blob_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseReadTensorboardBlobData, + "_BaseReadTensorboardBlobData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["timeSeries"] = 'time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).read_tensorboard_blob_data._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("blob_ids", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("blobIds", )) # verify required fields with non-default values are left alone assert "timeSeries" in jsonified_request @@ -16421,13 +16143,6 @@ def test_read_tensorboard_blob_data_rest_required_fields(request_type=tensorboar assert sorted(expected_params) == sorted(actual_params) -def test_read_tensorboard_blob_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.read_tensorboard_blob_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(("blobIds", )) & set(("timeSeries", ))) - - def test_read_tensorboard_blob_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16533,16 +16248,18 @@ def test_write_tensorboard_experiment_data_rest_required_fields(request_type=ten # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_experiment_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseWriteTensorboardExperimentData, + "_BaseWriteTensorboardExperimentData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardExperiment"] = 'tensorboard_experiment_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_experiment_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboardExperiment" in jsonified_request assert jsonified_request["tensorboardExperiment"] == 'tensorboard_experiment_value' @@ -16592,13 +16309,6 @@ def test_write_tensorboard_experiment_data_rest_required_fields(request_type=ten assert sorted(expected_params) == sorted(actual_params) -def test_write_tensorboard_experiment_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.write_tensorboard_experiment_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboardExperiment", "writeRunDataRequests", ))) - - def test_write_tensorboard_experiment_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16703,16 +16413,18 @@ def test_write_tensorboard_run_data_rest_required_fields(request_type=tensorboar # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_run_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseWriteTensorboardRunData, + "_BaseWriteTensorboardRunData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardRun"] = 'tensorboard_run_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).write_tensorboard_run_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboardRun" in jsonified_request assert jsonified_request["tensorboardRun"] == 'tensorboard_run_value' @@ -16762,13 +16474,6 @@ def test_write_tensorboard_run_data_rest_required_fields(request_type=tensorboar assert sorted(expected_params) == sorted(actual_params) -def test_write_tensorboard_run_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.write_tensorboard_run_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboardRun", "timeSeriesData", ))) - - def test_write_tensorboard_run_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -16873,16 +16578,18 @@ def test_export_tensorboard_time_series_data_rest_required_fields(request_type=t # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseExportTensorboardTimeSeriesData, + "_BaseExportTensorboardTimeSeriesData__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["tensorboardTimeSeries"] = 'tensorboard_time_series_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).export_tensorboard_time_series_data._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "tensorboardTimeSeries" in jsonified_request assert jsonified_request["tensorboardTimeSeries"] == 'tensorboard_time_series_value' @@ -16932,13 +16639,6 @@ def test_export_tensorboard_time_series_data_rest_required_fields(request_type=t assert sorted(expected_params) == sorted(actual_params) -def test_export_tensorboard_time_series_data_rest_unset_required_fields(): - transport = transports.TensorboardServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.export_tensorboard_time_series_data._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("tensorboardTimeSeries", ))) - - def test_export_tensorboard_time_series_data_rest_flattened(): client = TensorboardServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_data_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_data_service.py index a25f20e868..6a32470f9f 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_data_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_data_service.py @@ -91,6 +91,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -134,150 +146,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert VertexRagDataServiceClient._get_default_mtls_endpoint(None) is None - assert VertexRagDataServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert VertexRagDataServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert VertexRagDataServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagDataServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - VertexRagDataServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert VertexRagDataServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagDataServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert VertexRagDataServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - VertexRagDataServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert VertexRagDataServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -291,40 +159,6 @@ def test__get_client_cert_source(): assert VertexRagDataServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert VertexRagDataServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(VertexRagDataServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagDataServiceClient)) -@mock.patch.object(VertexRagDataServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagDataServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = VertexRagDataServiceClient._DEFAULT_UNIVERSE - default_endpoint = VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = VertexRagDataServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert VertexRagDataServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert VertexRagDataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagDataServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert VertexRagDataServiceClient._get_api_endpoint(None, None, default_universe, "always") == VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagDataServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == VertexRagDataServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagDataServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert VertexRagDataServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagDataServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert VertexRagDataServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert VertexRagDataServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert VertexRagDataServiceClient._get_universe_domain(None, None) == VertexRagDataServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - VertexRagDataServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -709,11 +543,12 @@ def test_vertex_rag_data_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -756,11 +591,12 @@ def test_vertex_rag_data_service_client_get_mtls_endpoint_and_cert_source(client for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -9428,16 +9264,18 @@ def test_create_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateRagCorpus, + "_BaseCreateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_corpus._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -9484,13 +9322,6 @@ def test_create_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser assert sorted(expected_params) == sorted(actual_params) -def test_create_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "ragCorpus", ))) - - def test_create_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9596,14 +9427,16 @@ def test_update_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateRagCorpus, + "_BaseUpdateRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_corpus._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone client = VertexRagDataServiceClient( @@ -9648,13 +9481,6 @@ def test_update_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser assert sorted(expected_params) == sorted(actual_params) -def test_update_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("ragCorpus", ))) - - def test_update_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9755,16 +9581,18 @@ def test_get_rag_corpus_rest_required_fields(request_type=vertex_rag_data_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagCorpus, + "_BaseGetRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_corpus._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -9813,13 +9641,6 @@ def test_get_rag_corpus_rest_required_fields(request_type=vertex_rag_data_servic assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -9922,17 +9743,20 @@ def test_list_rag_corpora_rest_required_fields(request_type=vertex_rag_data_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_corpora._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListRagCorpora, + "_BaseListRagCorpora__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_corpora._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -9982,13 +9806,6 @@ def test_list_rag_corpora_rest_required_fields(request_type=vertex_rag_data_serv assert sorted(expected_params) == sorted(actual_params) -def test_list_rag_corpora_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_rag_corpora._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_rag_corpora_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10160,17 +9977,20 @@ def test_delete_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_corpus._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteRagCorpus, + "_BaseDeleteRagCorpus__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_corpus._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("force", "force_delete", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("force", "forceDelete", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -10217,13 +10037,6 @@ def test_delete_rag_corpus_rest_required_fields(request_type=vertex_rag_data_ser assert sorted(expected_params) == sorted(actual_params) -def test_delete_rag_corpus_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_rag_corpus._get_unset_required_fields({}) - assert set(unset_fields) == (set(("force", "forceDelete", )) & set(("name", ))) - - def test_delete_rag_corpus_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10324,16 +10137,18 @@ def test_upload_rag_file_rest_required_fields(request_type=vertex_rag_data_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_rag_file._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUploadRagFile, + "_BaseUploadRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).upload_rag_file._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10383,13 +10198,6 @@ def test_upload_rag_file_rest_required_fields(request_type=vertex_rag_data_servi assert sorted(expected_params) == sorted(actual_params) -def test_upload_rag_file_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.upload_rag_file._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "ragFile", "uploadRagFileConfig", ))) - - def test_upload_rag_file_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10500,16 +10308,18 @@ def test_import_rag_files_rest_required_fields(request_type=vertex_rag_data_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_rag_files._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseImportRagFiles, + "_BaseImportRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).import_rag_files._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -10556,13 +10366,6 @@ def test_import_rag_files_rest_required_fields(request_type=vertex_rag_data_serv assert sorted(expected_params) == sorted(actual_params) -def test_import_rag_files_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.import_rag_files._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "importRagFilesConfig", ))) - - def test_import_rag_files_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10665,16 +10468,18 @@ def test_get_rag_file_rest_required_fields(request_type=vertex_rag_data_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_file._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagFile, + "_BaseGetRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_file._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -10723,13 +10528,6 @@ def test_get_rag_file_rest_required_fields(request_type=vertex_rag_data_service. assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_file_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_file._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_file_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -10832,17 +10630,20 @@ def test_list_rag_files_rest_required_fields(request_type=vertex_rag_data_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_files._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListRagFiles, + "_BaseListRagFiles__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_files._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -10892,13 +10693,6 @@ def test_list_rag_files_rest_required_fields(request_type=vertex_rag_data_servic assert sorted(expected_params) == sorted(actual_params) -def test_list_rag_files_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_rag_files._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_rag_files_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11070,17 +10864,20 @@ def test_delete_rag_file_rest_required_fields(request_type=vertex_rag_data_servi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_file._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteRagFile, + "_BaseDeleteRagFile__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_file._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("force_delete", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("forceDelete", )) # verify required fields with non-default values are left alone assert "name" in jsonified_request @@ -11127,13 +10924,6 @@ def test_delete_rag_file_rest_required_fields(request_type=vertex_rag_data_servi assert sorted(expected_params) == sorted(actual_params) -def test_delete_rag_file_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_rag_file._get_unset_required_fields({}) - assert set(unset_fields) == (set(("forceDelete", )) & set(("name", ))) - - def test_delete_rag_file_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11237,14 +11027,16 @@ def test_update_rag_engine_config_rest_required_fields(request_type=vertex_rag_d # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_engine_config._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateRagEngineConfig, + "_BaseUpdateRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_engine_config._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone client = VertexRagDataServiceClient( @@ -11289,13 +11081,6 @@ def test_update_rag_engine_config_rest_required_fields(request_type=vertex_rag_d assert sorted(expected_params) == sorted(actual_params) -def test_update_rag_engine_config_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_rag_engine_config._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("ragEngineConfig", ))) - - def test_update_rag_engine_config_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11396,16 +11181,18 @@ def test_get_rag_engine_config_rest_required_fields(request_type=vertex_rag_data # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_engine_config._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagEngineConfig, + "_BaseGetRagEngineConfig__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_engine_config._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11454,13 +11241,6 @@ def test_get_rag_engine_config_rest_required_fields(request_type=vertex_rag_data assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_engine_config_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_engine_config._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_engine_config_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11563,17 +11343,20 @@ def test_create_rag_data_schema_rest_required_fields(request_type=vertex_rag_dat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_data_schema._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateRagDataSchema, + "_BaseCreateRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_data_schema._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("rag_data_schema_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("ragDataSchemaId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -11624,13 +11407,6 @@ def test_create_rag_data_schema_rest_required_fields(request_type=vertex_rag_dat assert sorted(expected_params) == sorted(actual_params) -def test_create_rag_data_schema_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_rag_data_schema._get_unset_required_fields({}) - assert set(unset_fields) == (set(("ragDataSchemaId", )) & set(("parent", "ragDataSchema", ))) - - def test_create_rag_data_schema_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -11741,16 +11517,18 @@ def test_batch_create_rag_data_schemas_rest_required_fields(request_type=vertex_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_rag_data_schemas._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateRagDataSchemas, + "_BaseBatchCreateRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_rag_data_schemas._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -11797,13 +11575,6 @@ def test_batch_create_rag_data_schemas_rest_required_fields(request_type=vertex_ assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_rag_data_schemas_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_rag_data_schemas._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_get_rag_data_schema_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -11852,16 +11623,18 @@ def test_get_rag_data_schema_rest_required_fields(request_type=vertex_rag_data_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_data_schema._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagDataSchema, + "_BaseGetRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_data_schema._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -11910,13 +11683,6 @@ def test_get_rag_data_schema_rest_required_fields(request_type=vertex_rag_data_s assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_data_schema_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_data_schema._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_data_schema_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12019,17 +11785,20 @@ def test_list_rag_data_schemas_rest_required_fields(request_type=vertex_rag_data # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_data_schemas._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListRagDataSchemas, + "_BaseListRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_data_schemas._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12079,13 +11848,6 @@ def test_list_rag_data_schemas_rest_required_fields(request_type=vertex_rag_data assert sorted(expected_params) == sorted(actual_params) -def test_list_rag_data_schemas_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_rag_data_schemas._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_rag_data_schemas_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12253,16 +12015,18 @@ def test_delete_rag_data_schema_rest_required_fields(request_type=vertex_rag_dat # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_data_schema._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteRagDataSchema, + "_BaseDeleteRagDataSchema__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_data_schema._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -12308,13 +12072,6 @@ def test_delete_rag_data_schema_rest_required_fields(request_type=vertex_rag_dat assert sorted(expected_params) == sorted(actual_params) -def test_delete_rag_data_schema_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_rag_data_schema._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_rag_data_schema_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12420,7 +12177,12 @@ def test_batch_delete_rag_data_schemas_rest_required_fields(request_type=vertex_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_rag_data_schemas._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchDeleteRagDataSchemas, + "_BaseBatchDeleteRagDataSchemas__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -12428,9 +12190,6 @@ def test_batch_delete_rag_data_schemas_rest_required_fields(request_type=vertex_ jsonified_request["parent"] = 'parent_value' jsonified_request["names"] = 'names_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_rag_data_schemas._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -12479,13 +12238,6 @@ def test_batch_delete_rag_data_schemas_rest_required_fields(request_type=vertex_ assert sorted(expected_params) == sorted(actual_params) -def test_batch_delete_rag_data_schemas_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_delete_rag_data_schemas._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "names", ))) - - def test_create_rag_metadata_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -12534,17 +12286,20 @@ def test_create_rag_metadata_rest_required_fields(request_type=vertex_rag_data_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_metadata._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateRagMetadata, + "_BaseCreateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_rag_metadata._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("rag_metadata_id", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("ragMetadataId", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -12595,13 +12350,6 @@ def test_create_rag_metadata_rest_required_fields(request_type=vertex_rag_data_s assert sorted(expected_params) == sorted(actual_params) -def test_create_rag_metadata_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_rag_metadata._get_unset_required_fields({}) - assert set(unset_fields) == (set(("ragMetadataId", )) & set(("parent", "ragMetadata", ))) - - def test_create_rag_metadata_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12712,16 +12460,18 @@ def test_batch_create_rag_metadata_rest_required_fields(request_type=vertex_rag_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_rag_metadata._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchCreateRagMetadata, + "_BaseBatchCreateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_create_rag_metadata._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -12768,13 +12518,6 @@ def test_batch_create_rag_metadata_rest_required_fields(request_type=vertex_rag_ assert sorted(expected_params) == sorted(actual_params) -def test_batch_create_rag_metadata_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_create_rag_metadata._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "requests", ))) - - def test_update_rag_metadata_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -12822,14 +12565,16 @@ def test_update_rag_metadata_rest_required_fields(request_type=vertex_rag_data_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_metadata._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseUpdateRagMetadata, + "_BaseUpdateRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).update_rag_metadata._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone client = VertexRagDataServiceClient( @@ -12877,13 +12622,6 @@ def test_update_rag_metadata_rest_required_fields(request_type=vertex_rag_data_s assert sorted(expected_params) == sorted(actual_params) -def test_update_rag_metadata_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.update_rag_metadata._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("ragMetadata", ))) - - def test_update_rag_metadata_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -12986,16 +12724,18 @@ def test_get_rag_metadata_rest_required_fields(request_type=vertex_rag_data_serv # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_metadata._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetRagMetadata, + "_BaseGetRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_rag_metadata._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13044,13 +12784,6 @@ def test_get_rag_metadata_rest_required_fields(request_type=vertex_rag_data_serv assert sorted(expected_params) == sorted(actual_params) -def test_get_rag_metadata_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_rag_metadata._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_rag_metadata_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13153,17 +12886,20 @@ def test_list_rag_metadata_rest_required_fields(request_type=vertex_rag_data_ser # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_metadata._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListRagMetadata, + "_BaseListRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_rag_metadata._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -13213,13 +12949,6 @@ def test_list_rag_metadata_rest_required_fields(request_type=vertex_rag_data_ser assert sorted(expected_params) == sorted(actual_params) -def test_list_rag_metadata_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_rag_metadata._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_rag_metadata_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13387,16 +13116,18 @@ def test_delete_rag_metadata_rest_required_fields(request_type=vertex_rag_data_s # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_metadata._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteRagMetadata, + "_BaseDeleteRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_rag_metadata._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -13442,13 +13173,6 @@ def test_delete_rag_metadata_rest_required_fields(request_type=vertex_rag_data_s assert sorted(expected_params) == sorted(actual_params) -def test_delete_rag_metadata_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_rag_metadata._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_rag_metadata_rest_flattened(): client = VertexRagDataServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -13554,7 +13278,12 @@ def test_batch_delete_rag_metadata_rest_required_fields(request_type=vertex_rag_ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_rag_metadata._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseBatchDeleteRagMetadata, + "_BaseBatchDeleteRagMetadata__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -13562,9 +13291,6 @@ def test_batch_delete_rag_metadata_rest_required_fields(request_type=vertex_rag_ jsonified_request["parent"] = 'parent_value' jsonified_request["names"] = 'names_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).batch_delete_rag_metadata._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -13613,13 +13339,6 @@ def test_batch_delete_rag_metadata_rest_required_fields(request_type=vertex_rag_ assert sorted(expected_params) == sorted(actual_params) -def test_batch_delete_rag_metadata_rest_unset_required_fields(): - transport = transports.VertexRagDataServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.batch_delete_rag_metadata._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "names", ))) - - def test_credentials_transport_error(): # It is an error to provide credentials and a transport instance. transport = transports.VertexRagDataServiceGrpcTransport( diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_service.py index 8d8d2c6e06..947dbff04b 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_vertex_rag_service.py @@ -89,6 +89,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -132,150 +144,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert VertexRagServiceClient._get_default_mtls_endpoint(None) is None - assert VertexRagServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert VertexRagServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert VertexRagServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - VertexRagServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert VertexRagServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert VertexRagServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert VertexRagServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - VertexRagServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert VertexRagServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert VertexRagServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -289,40 +157,6 @@ def test__get_client_cert_source(): assert VertexRagServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert VertexRagServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(VertexRagServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagServiceClient)) -@mock.patch.object(VertexRagServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VertexRagServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = VertexRagServiceClient._DEFAULT_UNIVERSE - default_endpoint = VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = VertexRagServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert VertexRagServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert VertexRagServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert VertexRagServiceClient._get_api_endpoint(None, None, default_universe, "always") == VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == VertexRagServiceClient.DEFAULT_MTLS_ENDPOINT - assert VertexRagServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert VertexRagServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - VertexRagServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert VertexRagServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert VertexRagServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert VertexRagServiceClient._get_universe_domain(None, None) == VertexRagServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - VertexRagServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -707,11 +541,12 @@ def test_vertex_rag_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -754,11 +589,12 @@ def test_vertex_rag_service_client_get_mtls_endpoint_and_cert_source(client_clas for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -2650,16 +2486,18 @@ def test_retrieve_contexts_rest_required_fields(request_type=vertex_rag_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).retrieve_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseRetrieveContexts, + "_BaseRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).retrieve_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2709,13 +2547,6 @@ def test_retrieve_contexts_rest_required_fields(request_type=vertex_rag_service. assert sorted(expected_params) == sorted(actual_params) -def test_retrieve_contexts_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.retrieve_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "query", ))) - - def test_retrieve_contexts_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2820,16 +2651,18 @@ def test_augment_prompt_rest_required_fields(request_type=vertex_rag_service.Aug # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).augment_prompt._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAugmentPrompt, + "_BaseAugmentPrompt__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).augment_prompt._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -2879,13 +2712,6 @@ def test_augment_prompt_rest_required_fields(request_type=vertex_rag_service.Aug assert sorted(expected_params) == sorted(actual_params) -def test_augment_prompt_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.augment_prompt._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_augment_prompt_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -2991,16 +2817,18 @@ def test_corroborate_content_rest_required_fields(request_type=vertex_rag_servic # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).corroborate_content._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCorroborateContent, + "_BaseCorroborateContent__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).corroborate_content._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3050,13 +2878,6 @@ def test_corroborate_content_rest_required_fields(request_type=vertex_rag_servic assert sorted(expected_params) == sorted(actual_params) -def test_corroborate_content_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.corroborate_content._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_corroborate_content_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3163,16 +2984,18 @@ def test_ask_contexts_rest_required_fields(request_type=vertex_rag_service.AskCo # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).ask_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAskContexts, + "_BaseAskContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).ask_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3222,13 +3045,6 @@ def test_ask_contexts_rest_required_fields(request_type=vertex_rag_service.AskCo assert sorted(expected_params) == sorted(actual_params) -def test_ask_contexts_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.ask_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "query", ))) - - def test_ask_contexts_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -3337,16 +3153,18 @@ def test_async_retrieve_contexts_rest_required_fields(request_type=vertex_rag_se # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_retrieve_contexts._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAsyncRetrieveContexts, + "_BaseAsyncRetrieveContexts__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).async_retrieve_contexts._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -3393,13 +3211,6 @@ def test_async_retrieve_contexts_rest_required_fields(request_type=vertex_rag_se assert sorted(expected_params) == sorted(actual_params) -def test_async_retrieve_contexts_rest_unset_required_fields(): - transport = transports.VertexRagServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.async_retrieve_contexts._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "query", ))) - - def test_async_retrieve_contexts_rest_flattened(): client = VertexRagServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/aiplatform_v1beta1/test_vizier_service.py b/tests/unit/gapic/aiplatform_v1beta1/test_vizier_service.py index 91a9b1b454..fb92650f66 100644 --- a/tests/unit/gapic/aiplatform_v1beta1/test_vizier_service.py +++ b/tests/unit/gapic/aiplatform_v1beta1/test_vizier_service.py @@ -89,6 +89,18 @@ CRED_INFO_STRING = json.dumps(CRED_INFO_JSON) +@pytest.fixture(autouse=True) +def disable_mtls_env(): + with mock.patch.dict( + os.environ, + { + "GOOGLE_API_USE_CLIENT_CERTIFICATE": "false", + "CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE": "false", + }, + ): + yield + + async def mock_async_gen(data, chunk_size=1): for i in range(0, len(data)): # pragma: NO COVER chunk = data[i : i + chunk_size] @@ -132,150 +144,6 @@ def set_event_loop(): asyncio.set_event_loop(None) -def test__get_default_mtls_endpoint(): - api_endpoint = "example.googleapis.com" - api_mtls_endpoint = "example.mtls.googleapis.com" - sandbox_endpoint = "example.sandbox.googleapis.com" - sandbox_mtls_endpoint = "example.mtls.sandbox.googleapis.com" - non_googleapi = "api.example.com" - custom_endpoint = ".custom" - - assert VizierServiceClient._get_default_mtls_endpoint(None) is None - assert VizierServiceClient._get_default_mtls_endpoint(api_endpoint) == api_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(api_mtls_endpoint) == api_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(sandbox_endpoint) == sandbox_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(sandbox_mtls_endpoint) == sandbox_mtls_endpoint - assert VizierServiceClient._get_default_mtls_endpoint(non_googleapi) == non_googleapi - assert VizierServiceClient._get_default_mtls_endpoint(custom_endpoint) == custom_endpoint - -def test__read_environment_variables(): - assert VizierServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VizierServiceClient._read_environment_variables() == (True, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VizierServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict( - os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "Unsupported"} - ): - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with pytest.raises(ValueError) as excinfo: - VizierServiceClient._read_environment_variables() - assert ( - str(excinfo.value) - == "Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`" - ) - else: - assert VizierServiceClient._read_environment_variables() == ( - False, - "auto", - None, - ) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "never"}): - assert VizierServiceClient._read_environment_variables() == (False, "never", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "always"}): - assert VizierServiceClient._read_environment_variables() == (False, "always", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "auto"}): - assert VizierServiceClient._read_environment_variables() == (False, "auto", None) - - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "Unsupported"}): - with pytest.raises(MutualTLSChannelError) as excinfo: - VizierServiceClient._read_environment_variables() - assert str(excinfo.value) == "Environment variable `GOOGLE_API_USE_MTLS_ENDPOINT` must be `never`, `auto` or `always`" - - with mock.patch.dict(os.environ, {"GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): - assert VizierServiceClient._read_environment_variables() == (False, "auto", "foo.com") - - -def test_use_client_cert_effective(): - # Test case 1: Test when `should_use_client_cert` returns True. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=True): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 2: Test when `should_use_client_cert` returns False. - # We mock the `should_use_client_cert` function to simulate a scenario where - # the google-auth library supports automatic mTLS and determines that a - # client certificate should NOT be used. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch("google.auth.transport.mtls.should_use_client_cert", return_value=False): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 3: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "true". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 4: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 5: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "True". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "True"}): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 6: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "False". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "False"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 7: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "TRUE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "TRUE"}): - assert VizierServiceClient._use_client_cert_effective() is True - - # Test case 8: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to "FALSE". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "FALSE"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 9: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is not set. - # In this case, the method should return False, which is the default value. - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, clear=True): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 10: Test when `should_use_client_cert` is unavailable and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should raise a ValueError as the environment variable must be either - # "true" or "false". - if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - with pytest.raises(ValueError): - VizierServiceClient._use_client_cert_effective() - - # Test case 11: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is set to an invalid value. - # The method should return False as the environment variable is set to an invalid value. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "unsupported"}): - assert VizierServiceClient._use_client_cert_effective() is False - - # Test case 12: Test when `should_use_client_cert` is available and the - # `GOOGLE_API_USE_CLIENT_CERTIFICATE` environment variable is unset. Also, - # the GOOGLE_API_CONFIG environment variable is unset. - if hasattr(google.auth.transport.mtls, "should_use_client_cert"): - with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": ""}): - with mock.patch.dict(os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": ""}): - assert VizierServiceClient._use_client_cert_effective() is False - def test__get_client_cert_source(): mock_provided_cert_source = mock.Mock() mock_default_cert_source = mock.Mock() @@ -289,40 +157,6 @@ def test__get_client_cert_source(): assert VizierServiceClient._get_client_cert_source(None, True) is mock_default_cert_source assert VizierServiceClient._get_client_cert_source(mock_provided_cert_source, "true") is mock_provided_cert_source -@mock.patch.object(VizierServiceClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VizierServiceClient)) -@mock.patch.object(VizierServiceAsyncClient, "_DEFAULT_ENDPOINT_TEMPLATE", modify_default_endpoint_template(VizierServiceAsyncClient)) -def test__get_api_endpoint(): - api_override = "foo.com" - mock_client_cert_source = mock.Mock() - default_universe = VizierServiceClient._DEFAULT_UNIVERSE - default_endpoint = VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=default_universe) - mock_universe = "bar.com" - mock_endpoint = VizierServiceClient._DEFAULT_ENDPOINT_TEMPLATE.format(UNIVERSE_DOMAIN=mock_universe) - - assert VizierServiceClient._get_api_endpoint(api_override, mock_client_cert_source, default_universe, "always") == api_override - assert VizierServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "auto") == VizierServiceClient.DEFAULT_MTLS_ENDPOINT - assert VizierServiceClient._get_api_endpoint(None, None, default_universe, "auto") == default_endpoint - assert VizierServiceClient._get_api_endpoint(None, None, default_universe, "always") == VizierServiceClient.DEFAULT_MTLS_ENDPOINT - assert VizierServiceClient._get_api_endpoint(None, mock_client_cert_source, default_universe, "always") == VizierServiceClient.DEFAULT_MTLS_ENDPOINT - assert VizierServiceClient._get_api_endpoint(None, None, mock_universe, "never") == mock_endpoint - assert VizierServiceClient._get_api_endpoint(None, None, default_universe, "never") == default_endpoint - - with pytest.raises(MutualTLSChannelError) as excinfo: - VizierServiceClient._get_api_endpoint(None, mock_client_cert_source, mock_universe, "auto") - assert str(excinfo.value) == "mTLS is not supported in any universe other than googleapis.com." - - -def test__get_universe_domain(): - client_universe_domain = "foo.com" - universe_domain_env = "bar.com" - - assert VizierServiceClient._get_universe_domain(client_universe_domain, universe_domain_env) == client_universe_domain - assert VizierServiceClient._get_universe_domain(None, universe_domain_env) == universe_domain_env - assert VizierServiceClient._get_universe_domain(None, None) == VizierServiceClient._DEFAULT_UNIVERSE - - with pytest.raises(ValueError) as excinfo: - VizierServiceClient._get_universe_domain("", None) - assert str(excinfo.value) == "Universe Domain cannot be an empty string." @pytest.mark.parametrize("error_code,cred_info_json,show_cred_info", [ (401, CRED_INFO_JSON, True), @@ -707,11 +541,12 @@ def test_vizier_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", None) + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", None) with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -754,11 +589,12 @@ def test_vizier_service_client_get_mtls_endpoint_and_cert_source(client_class): for config_data, expected_cert_source in test_cases: env = os.environ.copy() env.pop("GOOGLE_API_USE_CLIENT_CERTIFICATE", "") + env.pop("CLOUDSDK_CONTEXT_AWARE_USE_CLIENT_CERTIFICATE", "") with mock.patch.dict(os.environ, env, clear=True): config_filename = "mock_certificate_config.json" config_file_content = json.dumps(config_data) m = mock.mock_open(read_data=config_file_content) - with mock.patch("builtins.open", m): + with mock.patch("builtins.open", m), mock.patch("os.path.exists", side_effect=lambda path: os.path.basename(path) == config_filename): with mock.patch.dict( os.environ, {"GOOGLE_API_CERTIFICATE_CONFIG": config_filename} ): @@ -5824,16 +5660,18 @@ def test_create_study_rest_required_fields(request_type=vizier_service.CreateStu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateStudy, + "_BaseCreateStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -5883,13 +5721,6 @@ def test_create_study_rest_required_fields(request_type=vizier_service.CreateStu assert sorted(expected_params) == sorted(actual_params) -def test_create_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "study", ))) - - def test_create_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -5994,16 +5825,18 @@ def test_get_study_rest_required_fields(request_type=vizier_service.GetStudyRequ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetStudy, + "_BaseGetStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6052,13 +5885,6 @@ def test_get_study_rest_required_fields(request_type=vizier_service.GetStudyRequ assert sorted(expected_params) == sorted(actual_params) -def test_get_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6161,17 +5987,20 @@ def test_list_studies_rest_required_fields(request_type=vizier_service.ListStudi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_studies._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListStudies, + "_BaseListStudies__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_studies._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -6221,13 +6050,6 @@ def test_list_studies_rest_required_fields(request_type=vizier_service.ListStudi assert sorted(expected_params) == sorted(actual_params) -def test_list_studies_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_studies._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_studies_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6395,16 +6217,18 @@ def test_delete_study_rest_required_fields(request_type=vizier_service.DeleteStu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteStudy, + "_BaseDeleteStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -6450,13 +6274,6 @@ def test_delete_study_rest_required_fields(request_type=vizier_service.DeleteStu assert sorted(expected_params) == sorted(actual_params) -def test_delete_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6558,7 +6375,12 @@ def test_lookup_study_rest_required_fields(request_type=vizier_service.LookupStu # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).lookup_study._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseLookupStudy, + "_BaseLookupStudy__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -6566,9 +6388,6 @@ def test_lookup_study_rest_required_fields(request_type=vizier_service.LookupStu jsonified_request["parent"] = 'parent_value' jsonified_request["displayName"] = 'display_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).lookup_study._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6620,13 +6439,6 @@ def test_lookup_study_rest_required_fields(request_type=vizier_service.LookupStu assert sorted(expected_params) == sorted(actual_params) -def test_lookup_study_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.lookup_study._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "displayName", ))) - - def test_lookup_study_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -6735,7 +6547,12 @@ def test_suggest_trials_rest_required_fields(request_type=vizier_service.Suggest # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).suggest_trials._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseSuggestTrials, + "_BaseSuggestTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present @@ -6744,9 +6561,6 @@ def test_suggest_trials_rest_required_fields(request_type=vizier_service.Suggest jsonified_request["suggestionCount"] = 1744 jsonified_request["clientId"] = 'client_id_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).suggest_trials._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6797,13 +6611,6 @@ def test_suggest_trials_rest_required_fields(request_type=vizier_service.Suggest assert sorted(expected_params) == sorted(actual_params) -def test_suggest_trials_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.suggest_trials._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "suggestionCount", "clientId", ))) - - def test_create_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -6852,16 +6659,18 @@ def test_create_trial_rest_required_fields(request_type=vizier_service.CreateTri # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCreateTrial, + "_BaseCreateTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).create_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -6911,13 +6720,6 @@ def test_create_trial_rest_required_fields(request_type=vizier_service.CreateTri assert sorted(expected_params) == sorted(actual_params) -def test_create_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.create_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", "trial", ))) - - def test_create_trial_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7022,16 +6824,18 @@ def test_get_trial_rest_required_fields(request_type=vizier_service.GetTrialRequ # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseGetTrial, + "_BaseGetTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).get_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7080,13 +6884,6 @@ def test_get_trial_rest_required_fields(request_type=vizier_service.GetTrialRequ assert sorted(expected_params) == sorted(actual_params) -def test_get_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.get_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_get_trial_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7189,17 +6986,20 @@ def test_list_trials_rest_required_fields(request_type=vizier_service.ListTrials # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_trials._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListTrials, + "_BaseListTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_trials._get_unset_required_fields(jsonified_request) # Check that path parameters and body parameters are not mixing in. - assert not set(unset_fields) - set(("page_size", "page_token", )) - jsonified_request.update(unset_fields) + assert not set(unset_fields) - set(("pageSize", "pageToken", )) # verify required fields with non-default values are left alone assert "parent" in jsonified_request @@ -7249,13 +7049,6 @@ def test_list_trials_rest_required_fields(request_type=vizier_service.ListTrials assert sorted(expected_params) == sorted(actual_params) -def test_list_trials_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_trials._get_unset_required_fields({}) - assert set(unset_fields) == (set(("pageSize", "pageToken", )) & set(("parent", ))) - - def test_list_trials_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7423,16 +7216,18 @@ def test_add_trial_measurement_rest_required_fields(request_type=vizier_service. # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_trial_measurement._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseAddTrialMeasurement, + "_BaseAddTrialMeasurement__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["trialName"] = 'trial_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).add_trial_measurement._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "trialName" in jsonified_request assert jsonified_request["trialName"] == 'trial_name_value' @@ -7482,13 +7277,6 @@ def test_add_trial_measurement_rest_required_fields(request_type=vizier_service. assert sorted(expected_params) == sorted(actual_params) -def test_add_trial_measurement_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.add_trial_measurement._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("trialName", "measurement", ))) - - def test_complete_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -7537,16 +7325,18 @@ def test_complete_trial_rest_required_fields(request_type=vizier_service.Complet # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).complete_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCompleteTrial, + "_BaseCompleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).complete_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7596,13 +7386,6 @@ def test_complete_trial_rest_required_fields(request_type=vizier_service.Complet assert sorted(expected_params) == sorted(actual_params) -def test_complete_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.complete_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -7651,16 +7434,18 @@ def test_delete_trial_rest_required_fields(request_type=vizier_service.DeleteTri # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseDeleteTrial, + "_BaseDeleteTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).delete_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7706,13 +7491,6 @@ def test_delete_trial_rest_required_fields(request_type=vizier_service.DeleteTri assert sorted(expected_params) == sorted(actual_params) -def test_delete_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.delete_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_delete_trial_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), @@ -7817,16 +7595,18 @@ def test_check_trial_early_stopping_state_rest_required_fields(request_type=vizi # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).check_trial_early_stopping_state._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseCheckTrialEarlyStoppingState, + "_BaseCheckTrialEarlyStoppingState__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["trialName"] = 'trial_name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).check_trial_early_stopping_state._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "trialName" in jsonified_request assert jsonified_request["trialName"] == 'trial_name_value' @@ -7873,13 +7653,6 @@ def test_check_trial_early_stopping_state_rest_required_fields(request_type=vizi assert sorted(expected_params) == sorted(actual_params) -def test_check_trial_early_stopping_state_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.check_trial_early_stopping_state._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("trialName", ))) - - def test_stop_trial_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -7928,16 +7701,18 @@ def test_stop_trial_rest_required_fields(request_type=vizier_service.StopTrialRe # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_trial._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseStopTrial, + "_BaseStopTrial__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["name"] = 'name_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).stop_trial._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "name" in jsonified_request assert jsonified_request["name"] == 'name_value' @@ -7987,13 +7762,6 @@ def test_stop_trial_rest_required_fields(request_type=vizier_service.StopTrialRe assert sorted(expected_params) == sorted(actual_params) -def test_stop_trial_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.stop_trial._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("name", ))) - - def test_list_optimal_trials_rest_use_cached_wrapped_rpc(): # Clients should use _prep_wrapped_messages to create cached wrapped rpcs, # instead of constructing them on each call @@ -8042,16 +7810,18 @@ def test_list_optimal_trials_rest_required_fields(request_type=vizier_service.Li # verify fields with default values are dropped - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_optimal_trials._get_unset_required_fields(jsonified_request) + default_values = getattr( + transport_class._BaseListOptimalTrials, + "_BaseListOptimalTrials__REQUIRED_FIELDS_DEFAULT_VALUES", + {}, + ) + unset_fields = {k: v for k, v in default_values.items() if k not in jsonified_request} jsonified_request.update(unset_fields) # verify required fields with default values are now present jsonified_request["parent"] = 'parent_value' - unset_fields = transport_class(credentials=ga_credentials.AnonymousCredentials()).list_optimal_trials._get_unset_required_fields(jsonified_request) - jsonified_request.update(unset_fields) - # verify required fields with non-default values are left alone assert "parent" in jsonified_request assert jsonified_request["parent"] == 'parent_value' @@ -8101,13 +7871,6 @@ def test_list_optimal_trials_rest_required_fields(request_type=vizier_service.Li assert sorted(expected_params) == sorted(actual_params) -def test_list_optimal_trials_rest_unset_required_fields(): - transport = transports.VizierServiceRestTransport(credentials=ga_credentials.AnonymousCredentials) - - unset_fields = transport.list_optimal_trials._get_unset_required_fields({}) - assert set(unset_fields) == (set(()) & set(("parent", ))) - - def test_list_optimal_trials_rest_flattened(): client = VizierServiceClient( credentials=ga_credentials.AnonymousCredentials(), diff --git a/tests/unit/gapic/definition_v1/test_compat.py b/tests/unit/gapic/definition_v1/test_compat.py new file mode 100644 index 0000000000..c7fac90dbf --- /dev/null +++ b/tests/unit/gapic/definition_v1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1.schema.trainingjob.definition_v1._compat import transcode_request +from google.cloud.aiplatform.v1.schema.trainingjob.definition_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/definition_v1beta1/test_compat.py b/tests/unit/gapic/definition_v1beta1/test_compat.py new file mode 100644 index 0000000000..39d9dd13cd --- /dev/null +++ b/tests/unit/gapic/definition_v1beta1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1beta1.schema.trainingjob.definition_v1beta1._compat import transcode_request +from google.cloud.aiplatform.v1beta1.schema.trainingjob.definition_v1beta1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/instance_v1/test_compat.py b/tests/unit/gapic/instance_v1/test_compat.py new file mode 100644 index 0000000000..2c670b00ac --- /dev/null +++ b/tests/unit/gapic/instance_v1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1.schema.predict.instance_v1._compat import transcode_request +from google.cloud.aiplatform.v1.schema.predict.instance_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/instance_v1beta1/test_compat.py b/tests/unit/gapic/instance_v1beta1/test_compat.py new file mode 100644 index 0000000000..a6e2164ac5 --- /dev/null +++ b/tests/unit/gapic/instance_v1beta1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1beta1.schema.predict.instance_v1beta1._compat import transcode_request +from google.cloud.aiplatform.v1beta1.schema.predict.instance_v1beta1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/params_v1/test_compat.py b/tests/unit/gapic/params_v1/test_compat.py new file mode 100644 index 0000000000..4b44948ef0 --- /dev/null +++ b/tests/unit/gapic/params_v1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1.schema.predict.params_v1._compat import transcode_request +from google.cloud.aiplatform.v1.schema.predict.params_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/params_v1beta1/test_compat.py b/tests/unit/gapic/params_v1beta1/test_compat.py new file mode 100644 index 0000000000..1007cac694 --- /dev/null +++ b/tests/unit/gapic/params_v1beta1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1beta1.schema.predict.params_v1beta1._compat import transcode_request +from google.cloud.aiplatform.v1beta1.schema.predict.params_v1beta1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/prediction_v1/test_compat.py b/tests/unit/gapic/prediction_v1/test_compat.py new file mode 100644 index 0000000000..00a50304b8 --- /dev/null +++ b/tests/unit/gapic/prediction_v1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1.schema.predict.prediction_v1._compat import transcode_request +from google.cloud.aiplatform.v1.schema.predict.prediction_v1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables() diff --git a/tests/unit/gapic/prediction_v1beta1/test_compat.py b/tests/unit/gapic/prediction_v1beta1/test_compat.py new file mode 100644 index 0000000000..aaf8b6416b --- /dev/null +++ b/tests/unit/gapic/prediction_v1beta1/test_compat.py @@ -0,0 +1,425 @@ +# -*- coding: utf-8 -*- +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +"""Tests for the compatibility module for older versions of google-api-core.""" + +import json +import pytest +import os + +from unittest import mock + +import google.auth.transport.mtls + +from google.cloud.aiplatform.v1beta1.schema.predict.prediction_v1beta1._compat import transcode_request +from google.cloud.aiplatform.v1beta1.schema.predict.prediction_v1beta1._compat import get_universe_domain, get_api_endpoint, get_default_mtls_endpoint, should_use_client_cert, read_environment_variables + +from google.auth.exceptions import MutualTLSChannelError +from google.api_core.universe import EmptyUniverseError +from google.protobuf import descriptor_pb2 + + +def test_get_universe_domain(): + # When universe_domain is provided + assert ( + get_universe_domain("foo.com", default_universe="default.com") + == "foo.com" + ) + assert ( + get_universe_domain(" foo.com ", default_universe="default.com") + == "foo.com" + ) + + # When universe_domain is None, falls back to default_universe + assert ( + get_universe_domain(None, default_universe="default.com") + == "default.com" + ) + + # When multiple potential universes are provided, resolves in order of preference + assert ( + get_universe_domain( + "foo.com", "bar.com", default_universe="default.com" + ) + == "foo.com" + ) + assert ( + get_universe_domain(None, "bar.com", default_universe="default.com") + == "bar.com" + ) + assert ( + get_universe_domain(None, None, default_universe="default.com") + == "default.com" + ) + + # EmptyUniverseError raised when resolved value is empty string + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain("", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(" ", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + with pytest.raises(EmptyUniverseError) as excinfo: + get_universe_domain(None, "", default_universe="default.com") + assert str(excinfo.value) == "Universe Domain cannot be an empty string." + + +def test_get_default_mtls_endpoint(): + # Test valid API endpoints + assert ( + get_default_mtls_endpoint("foo.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com") + == "foo.mtls.sandbox.googleapis.com" + ) + # Test case-insensitivity + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com") + == "foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com") + == "foo.mtls.sandbox.googleapis.com" + ) + + # Test valid API endpoints with schemes + assert ( + get_default_mtls_endpoint("https://foo.googleapis.com") + == "https://foo.mtls.googleapis.com" + ) + assert ( + get_default_mtls_endpoint("http://foo.googleapis.com:8080/v1") + == "http://foo.mtls.googleapis.com:8080/v1" + ) + + # Test valid API endpoints with ports + assert ( + get_default_mtls_endpoint("foo.googleapis.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.sandbox.googleapis.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + # Test case-insensitivity with ports + assert ( + get_default_mtls_endpoint("foo.GoogleAPIs.com:443") + == "foo.mtls.googleapis.com:443" + ) + assert ( + get_default_mtls_endpoint("foo.Sandbox.GoogleAPIs.com:443") + == "foo.mtls.sandbox.googleapis.com:443" + ) + + # Test endpoints that shouldn't be converted + assert ( + get_default_mtls_endpoint("foo.mtls.googleapis.com") + == "foo.mtls.googleapis.com" + ) + assert get_default_mtls_endpoint("foo.com") == "foo.com" + assert get_default_mtls_endpoint("foo.com:8080") == "foo.com:8080" + + # Test empty/None endpoints + assert get_default_mtls_endpoint("") == "" + assert get_default_mtls_endpoint(None) is None + + # Test endpoints without host + assert get_default_mtls_endpoint("http://") == "http://" + assert get_default_mtls_endpoint("https://") == "https://" + + +@pytest.mark.parametrize( + "api_override,universe_domain,default_universe,default_mtls_endpoint,default_endpoint_template,use_mtls,expected", + [ + ( + "foo.com", + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + "foo.mtls.googleapis.com", + ), + ( + None, + "googleapis.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + False, + "foo.googleapis.com", + ), + ( + None, + "bar.com", + "googleapis.com", + "foo.mtls.googleapis.com", + "foo.{UNIVERSE_DOMAIN}", + True, + MutualTLSChannelError, + ), + ( + None, + "googleapis.com", + "googleapis.com", + None, + "foo.{UNIVERSE_DOMAIN}", + True, + ValueError, + ), + ], +) +def test_get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + expected, +): + if isinstance(expected, type) and issubclass(expected, Exception): + with pytest.raises(expected): + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + else: + assert ( + get_api_endpoint( + api_override, + universe_domain, + default_universe, + default_mtls_endpoint, + default_endpoint_template, + use_mtls, + ) + == expected + ) + + +def test_should_use_client_cert_fallback_env(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true"}, clear=True): + assert should_use_client_cert() is True + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "false"}, clear=True): + assert should_use_client_cert() is False + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "invalid"}, clear=True): + if not hasattr(google.auth.transport.mtls, "should_use_client_cert"): + with pytest.raises(ValueError, match="Environment variable `GOOGLE_API_USE_CLIENT_CERTIFICATE` must be either `true` or `false`"): + should_use_client_cert() + + +def test_transcode_basic(): + # We use FieldDescriptorProto as it has standard primitive fields and nested messages. + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is None + # 'number' should be in query parameters + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_nested_field(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{options.deprecated}/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "get" + assert transcoded["uri"] == "/v1/test/True/my-field" + assert body is None + assert "number" in query_params + assert query_params["number"] == 123 + + +def test_transcode_with_body(): + http_options = [ + { + "method": "post", + "uri": "/v1/test/{name}", + "body": "options", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.options.deprecated = True + request.number = 123 + + transcoded, body, query_params = transcode_request(http_options, request) + + assert transcoded["method"] == "post" + assert transcoded["uri"] == "/v1/test/my-field" + assert body is not None + body_data = json.loads(body) + assert body_data["deprecated"] is True + # Query parameters should not contain 'options' (the body) + assert "number" in query_params + assert query_params["number"] == 123 + assert "options" not in query_params + + +def test_transcode_with_required_fields_default_values(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + + required_defaults = {"requiredQueryParam": "default-val"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["requiredQueryParam"] == "default-val" + + +def test_transcode_with_numeric_enums(): + http_options = [ + { + "method": "get", + "uri": "/v1/test/{name}", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "my-field" + request.type = descriptor_pb2.FieldDescriptorProto.TYPE_STRING + + # Without numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=False + ) + assert query_params["type"] == "TYPE_STRING" + + # With numeric enums + _, _, query_params = transcode_request( + http_options, request, rest_numeric_enums=True + ) + # Type number for TYPE_STRING is 9 + assert query_params["type"] == 9 + assert query_params["$alt"] == "json;enum-encoding=int" + + +def test_transcode_no_query_params(): + http_options = [{"method": "get", "uri": "/v1/test"}] + request = descriptor_pb2.FieldDescriptorProto() + + with mock.patch( + "google.api_core.path_template.transcode", + return_value={"method": "get", "uri": "/v1/test"}, + ): + transcoded, body, query_params = transcode_request(http_options, request) + assert query_params == {} + + +def test_transcode_with_required_fields_existing_key(): + http_options = [ + { + "method": "get", + "uri": "/v1/test", + } + ] + + request = descriptor_pb2.FieldDescriptorProto() + request.name = "custom-name" + + required_defaults = {"name": "default-name"} + + transcoded, body, query_params = transcode_request( + http_options, + request, + required_fields_default_values=required_defaults, + ) + + assert query_params["name"] == "custom-name" + + +def test_transcode_request_invalid_request(): + http_options = [{"method": "get", "uri": "/v1/test"}] + with pytest.raises(TypeError, match="request cannot be None"): + transcode_request(http_options, None) + + +def test_transcode_request_proto_plus_wrapper(): + http_options = [{"method": "get", "uri": "/v1/test/{name}"}] + mock_pb = descriptor_pb2.FieldDescriptorProto() + mock_pb.name = "proto-plus-field" + + mock_proto_plus = mock.Mock() + mock_proto_plus._pb = mock_pb + + transcoded, _, _ = transcode_request(http_options, mock_proto_plus) + assert transcoded["uri"] == "/v1/test/proto-plus-field" + + +def test_read_environment_variables(): + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_CLIENT_CERTIFICATE": "true", "GOOGLE_API_USE_MTLS_ENDPOINT": "always", "GOOGLE_CLOUD_UNIVERSE_DOMAIN": "foo.com"}): + use_cert, mtls_endpoint, universe_domain = read_environment_variables() + assert use_cert is True + assert mtls_endpoint == "always" + assert universe_domain == "foo.com" + + with mock.patch.dict(os.environ, {"GOOGLE_API_USE_MTLS_ENDPOINT": "invalid"}): + with pytest.raises(MutualTLSChannelError): + read_environment_variables()