diff --git a/packages/google-cloud-spanner/google/cloud/spanner_dbapi/connection.py b/packages/google-cloud-spanner/google/cloud/spanner_dbapi/connection.py index a1570d9ca69a..c3539e9846c5 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_dbapi/connection.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_dbapi/connection.py @@ -820,6 +820,8 @@ def connect( instance_type=None, data_boost_enabled=False, auto_partition_mode=False, + username=None, + password=None, **kwargs, ): """Creates a connection to a Google Cloud Spanner database. @@ -909,6 +911,10 @@ def connect( :param client_key: (Optional) The path to the client key file used for mTLS connection. This is intended only for Spanner Omni endpoints. This is mandatory if Spanner Omni requires an mTLS connection. + :type username: str + :param username: (Optional) Username for Spanner Omni authentication. + :type password: str + :param password: (Optional) Password for Spanner Omni authentication. """ if client is None: client_info = ClientInfo( @@ -956,7 +962,29 @@ def connect( ) project = "default" - credentials = AnonymousCredentials() + has_username = username is not None + has_password = password is not None + if has_username != has_password: + raise ValueError( + "Both username and password must be specified for Omni authentication" + ) + from google.cloud.spanner_v1.omni.credentials import ( + SpannerOmniCredentials, + ) + + if has_username and has_password: + credentials = SpannerOmniCredentials( + username=username, + password=password, + target=host_endpoint, + use_plain_text=use_plain_text, + ca_certificate=ca_certificate, + client_certificate=client_certificate, + client_key=client_key, + ) + else: + credentials = AnonymousCredentials() + client_options = kwargs.get("client_options") if client_options is None: client_options = ClientOptions(api_endpoint=host_endpoint) @@ -967,8 +995,6 @@ def connect( import copy client_options = copy.copy(client_options) - client_options.api_endpoint = host_endpoint - client = spanner.Client( project=project, credentials=credentials, diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py index e02c79c6c553..6e256a4e24a5 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py @@ -75,6 +75,7 @@ def _create_spanner_omni_transport( client_certificate, client_key, interceptors=None, + credentials=None, ): """Creates a Spanner Omni transport in async mode. @@ -87,6 +88,8 @@ def _create_spanner_omni_transport( client_certificate (str): Path to the client certificate file for mTLS. client_key (str): Path to the client key file for mTLS. interceptors (list): Optional list of interceptors to add to the channel. + credentials (google.auth.credentials.Credentials, optional): Credentials + to use for authentication. Returns: object: An instance of the transport class created by `transport_factory`. @@ -98,8 +101,25 @@ def _create_spanner_omni_transport( from google.auth.credentials import AnonymousCredentials channel = None + all_interceptors = list(interceptors) if interceptors is not None else [] + if credentials is not None: + if hasattr(credentials, "create_async_auth_interceptors"): + all_interceptors.extend(credentials.create_async_auth_interceptors()) + elif hasattr(credentials, "create_async_auth_interceptor"): + res = credentials.create_async_auth_interceptor() + if isinstance(res, (list, tuple)): + all_interceptors.extend(res) + else: + all_interceptors.append(res) + elif hasattr(credentials, "create_auth_interceptor"): + res = credentials.create_auth_interceptor(is_async=True) + if isinstance(res, (list, tuple)): + all_interceptors.extend(res) + else: + all_interceptors.append(res) + if use_plain_text: - channel = grpc.aio.insecure_channel(target=host, interceptors=interceptors) + channel = grpc.aio.insecure_channel(target=host, interceptors=all_interceptors) elif ca_certificate: with open(ca_certificate, "rb") as f: ca_cert = f.read() @@ -119,12 +139,17 @@ def _create_spanner_omni_transport( ) else: ssl_creds = grpc.ssl_channel_credentials(root_certificates=ca_cert) - channel = grpc.aio.secure_channel(host, ssl_creds, interceptors=interceptors) + channel = grpc.aio.secure_channel( + host, ssl_creds, interceptors=all_interceptors + ) else: raise ValueError( "TLS/mTLS connection requires ca_certificate to be set for Spanner Omni" ) - return transport_factory(channel=channel, credentials=AnonymousCredentials()) + actual_credentials = ( + credentials if credentials is not None else AnonymousCredentials() + ) + return transport_factory(channel=channel, credentials=actual_credentials) def _create_experimental_host_transport( diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/client.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/client.py index 2e7da9fe808a..a18bbc9c605a 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/client.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/client.py @@ -302,6 +302,8 @@ def __init__( client_certificate=None, client_key=None, instance_type=None, + username=None, + password=None, ): self._emulator_host = _get_spanner_emulator_host() self._use_plain_text = use_plain_text @@ -353,10 +355,37 @@ def __init__( self._ca_certificate = ca_certificate self._client_certificate = client_certificate self._client_key = client_key - credentials = AnonymousCredentials() + self._host = host_endpoint + has_username = username is not None + has_password = password is not None + if has_username != has_password: + raise ValueError( + "Both username and password must be specified for Omni authentication" + ) + from google.cloud.spanner_v1.omni.credentials import ( + SpannerOmniCredentials, + ) + + if has_username and has_password: + credentials = SpannerOmniCredentials( + username=username, + password=password, + target=host_endpoint, + use_plain_text=use_plain_text, + ca_certificate=ca_certificate, + client_certificate=client_certificate, + client_key=client_key, + ) + elif not isinstance(credentials, SpannerOmniCredentials): + credentials = AnonymousCredentials() disable_builtin_metrics = True elif isinstance(credentials, AnonymousCredentials): self._emulator_host = self._client_options.api_endpoint + else: + if username is not None or password is not None: + raise ValueError( + "username and password can only be used when instance_type='omni'." + ) # NOTE: This API has no use for the _http argument, but sending it # will have no impact since the _http() @property only lazily @@ -509,6 +538,7 @@ def instance_admin_api(self): self._ca_certificate, self._client_certificate, self._client_key, + credentials=self.credentials, ) else: @@ -519,6 +549,7 @@ def instance_admin_api(self): self._ca_certificate, self._client_certificate, self._client_key, + credentials=self.credentials, ) self._instance_admin_api = InstanceAdminClient( @@ -567,6 +598,7 @@ def database_admin_api(self): self._ca_certificate, self._client_certificate, self._client_key, + credentials=self.credentials, ) else: @@ -577,6 +609,7 @@ def database_admin_api(self): self._ca_certificate, self._client_certificate, self._client_key, + credentials=self.credentials, ) self._database_admin_api = DatabaseAdminClient( diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py index e12c639eca3b..c12769172c23 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py @@ -513,6 +513,7 @@ def spanner_api(self): client._ca_certificate, client._client_certificate, client._client_key, + credentials=client.credentials, ) else: transport = _create_spanner_omni_transport_sync( @@ -522,6 +523,7 @@ def spanner_api(self): client._ca_certificate, client._client_certificate, client._client_key, + credentials=client.credentials, ) self._spanner_api = SpannerClient( client_info=client_info, diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/testing/database_test.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/testing/database_test.py index 59e99c51419c..d2c433eba4c9 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/testing/database_test.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/testing/database_test.py @@ -138,6 +138,7 @@ def spanner_api(self): client._client_certificate, client._client_key, self._interceptors, + credentials=client.credentials, ) else: transport = _create_spanner_omni_transport_sync( @@ -148,6 +149,7 @@ def spanner_api(self): client._client_certificate, client._client_key, self._interceptors, + credentials=client.credentials, ) self._spanner_api = SpannerClient( client_info=client_info, diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py index 06b137db1e28..a9740883afa4 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py @@ -1044,6 +1044,7 @@ def _create_spanner_omni_transport( client_certificate, client_key, interceptors=None, + credentials=None, ): """Creates a Spanner Omni transport. @@ -1056,6 +1057,8 @@ def _create_spanner_omni_transport( client_certificate (str): Path to the client certificate file for mTLS. client_key (str): Path to the client key file for mTLS. interceptors (list): Optional list of interceptors to add to the channel. + credentials (google.auth.credentials.Credentials, optional): Credentials + to use for authentication (e.g. `SpannerOmniCredentials`). Returns: object: An instance of the transport class created by `transport_factory`. @@ -1067,6 +1070,10 @@ def _create_spanner_omni_transport( from google.auth.credentials import AnonymousCredentials channel = None + all_interceptors = list(interceptors) if interceptors is not None else [] + if credentials is not None and hasattr(credentials, "create_auth_interceptor"): + all_interceptors.append(credentials.create_auth_interceptor()) + if use_plain_text: channel = grpc.insecure_channel(target=host) elif ca_certificate: @@ -1093,9 +1100,12 @@ def _create_spanner_omni_transport( raise ValueError( "TLS/mTLS connection requires ca_certificate to be set for Spanner Omni" ) - if interceptors is not None: - channel = grpc.intercept_channel(channel, *interceptors) - return transport_factory(channel=channel, credentials=AnonymousCredentials()) + if all_interceptors: + channel = grpc.intercept_channel(channel, *all_interceptors) + actual_credentials = ( + credentials if credentials is not None else AnonymousCredentials() + ) + return transport_factory(channel=channel, credentials=actual_credentials) def _create_experimental_host_transport( diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/client.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/client.py index 83bde6d89ccc..2ea9f3a85a07 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/client.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/client.py @@ -267,6 +267,8 @@ def __init__( client_certificate=None, client_key=None, instance_type=None, + username=None, + password=None, ): self._emulator_host = _get_spanner_emulator_host() self._use_plain_text = use_plain_text @@ -316,10 +318,37 @@ def __init__( self._ca_certificate = ca_certificate self._client_certificate = client_certificate self._client_key = client_key - credentials = AnonymousCredentials() + self._host = host_endpoint + has_username = username is not None + has_password = password is not None + if has_username != has_password: + raise ValueError( + "Both username and password must be specified for Omni authentication" + ) + from google.cloud.spanner_v1.omni.credentials import ( + SpannerOmniCredentials, + ) + + if has_username and has_password: + credentials = SpannerOmniCredentials( + username=username, + password=password, + target=host_endpoint, + use_plain_text=use_plain_text, + ca_certificate=ca_certificate, + client_certificate=client_certificate, + client_key=client_key, + ) + elif not isinstance(credentials, SpannerOmniCredentials): + credentials = AnonymousCredentials() disable_builtin_metrics = True elif isinstance(credentials, AnonymousCredentials): self._emulator_host = self._client_options.api_endpoint + else: + if username is not None or password is not None: + raise ValueError( + "username and password can only be used when instance_type='omni'." + ) super(Client, self).__init__( project=project, credentials=credentials, @@ -443,6 +472,7 @@ def instance_admin_api(self): self._ca_certificate, self._client_certificate, self._client_key, + credentials=self.credentials, ) self._instance_admin_api = InstanceAdminClient( client_info=self._client_info, @@ -481,6 +511,7 @@ def database_admin_api(self): self._ca_certificate, self._client_certificate, self._client_key, + credentials=self.credentials, ) self._database_admin_api = DatabaseAdminClient( client_info=self._client_info, diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py index f98d6d09f70d..37ae1e9f1267 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py @@ -452,6 +452,7 @@ def spanner_api(self): client._ca_certificate, client._client_certificate, client._client_key, + credentials=client.credentials, ) self._spanner_api = SpannerClient( client_info=client_info, diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/__init__.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/__init__.py new file mode 100644 index 000000000000..d295009a12c4 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/__init__.py @@ -0,0 +1,27 @@ +# -*- 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. +# + +"""Spanner Omni authentication and connection utilities.""" + +from google.cloud.spanner_v1.omni.credentials import SpannerOmniCredentials +from google.cloud.spanner_v1.omni.login_client import LoginClient +from google.cloud.spanner_v1.omni.opaque import UserAuthenticator + +__all__ = ( + "LoginClient", + "SpannerOmniCredentials", + "UserAuthenticator", +) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/credentials.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/credentials.py new file mode 100644 index 000000000000..3c966d6821d5 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/credentials.py @@ -0,0 +1,404 @@ +# -*- 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. +# + +"""Credentials implementation for Spanner Omni using OPAQUE login authentication.""" + +from __future__ import annotations + +import asyncio +import base64 +import datetime +import logging +import threading +from collections import namedtuple +from typing import Any, Callable, MutableMapping, Optional, Sequence + +import google.auth.credentials +import grpc +import grpc.aio + +from google.cloud.spanner_v1.omni.login_client import LoginClient + +_LOGGER = logging.getLogger(__name__) + + +class _ClientCallDetails( + namedtuple( + "_ClientCallDetails", + ["method", "timeout", "metadata", "credentials", "wait_for_ready"], + ), + grpc.ClientCallDetails, +): + pass + + +class _OmniAuthInterceptor( + grpc.UnaryUnaryClientInterceptor, + grpc.UnaryStreamClientInterceptor, + grpc.StreamUnaryClientInterceptor, + grpc.StreamStreamClientInterceptor, +): + """gRPC client interceptor that automatically attaches Spanner Omni Bearer tokens.""" + + def __init__(self, credentials: SpannerOmniCredentials) -> None: + self._credentials = credentials + + def _add_metadata( + self, client_call_details: grpc.ClientCallDetails + ) -> grpc.ClientCallDetails: + if not self._credentials.valid: + self._credentials.refresh() + token = self._credentials.token + metadata = list(client_call_details.metadata or []) + metadata.append(("authorization", f"Bearer {token}")) + + return _ClientCallDetails( + method=client_call_details.method, + timeout=client_call_details.timeout, + metadata=metadata, + credentials=client_call_details.credentials, + wait_for_ready=getattr(client_call_details, "wait_for_ready", None), + ) + + def intercept_unary_unary( + self, + continuation: Callable, + client_call_details: grpc.ClientCallDetails, + request: Any, + ) -> Any: + return continuation(self._add_metadata(client_call_details), request) + + def intercept_unary_stream( + self, + continuation: Callable, + client_call_details: grpc.ClientCallDetails, + request: Any, + ) -> Any: + return continuation(self._add_metadata(client_call_details), request) + + def intercept_stream_unary( + self, + continuation: Callable, + client_call_details: grpc.ClientCallDetails, + request_iterator: Any, + ) -> Any: + return continuation(self._add_metadata(client_call_details), request_iterator) + + def intercept_stream_stream( + self, + continuation: Callable, + client_call_details: grpc.ClientCallDetails, + request_iterator: Any, + ) -> Any: + return continuation(self._add_metadata(client_call_details), request_iterator) + + +class _AsyncClientCallDetails( + namedtuple( + "_AsyncClientCallDetails", + ["method", "timeout", "metadata", "credentials", "wait_for_ready"], + ), + grpc.aio.ClientCallDetails, +): + pass + + +class _AsyncBaseAuthInterceptor: + """Base helper for async auth interceptors.""" + + def __init__(self, credentials: SpannerOmniCredentials) -> None: + self._credentials = credentials + + async def _add_metadata( + self, client_call_details: grpc.aio.ClientCallDetails + ) -> grpc.aio.ClientCallDetails: + if not self._credentials.valid: + loop = asyncio.get_running_loop() + await loop.run_in_executor(None, self._credentials.refresh) + token = self._credentials.token + metadata = list(client_call_details.metadata or []) + metadata.append(("authorization", f"Bearer {token}")) + + return _AsyncClientCallDetails( + method=client_call_details.method, + timeout=client_call_details.timeout, + metadata=metadata, + credentials=client_call_details.credentials, + wait_for_ready=getattr(client_call_details, "wait_for_ready", None), + ) + + +class _AsyncUnaryUnaryAuthInterceptor( + _AsyncBaseAuthInterceptor, grpc.aio.UnaryUnaryClientInterceptor +): + """Async gRPC interceptor for unary-unary calls.""" + + async def intercept_unary_unary( + self, + continuation: Callable, + client_call_details: grpc.aio.ClientCallDetails, + request: Any, + ) -> Any: + return await continuation( + await self._add_metadata(client_call_details), request + ) + + +class _AsyncUnaryStreamAuthInterceptor( + _AsyncBaseAuthInterceptor, grpc.aio.UnaryStreamClientInterceptor +): + """Async gRPC interceptor for unary-stream calls.""" + + async def intercept_unary_stream( + self, + continuation: Callable, + client_call_details: grpc.aio.ClientCallDetails, + request: Any, + ) -> Any: + return await continuation( + await self._add_metadata(client_call_details), request + ) + + +class _AsyncStreamUnaryAuthInterceptor( + _AsyncBaseAuthInterceptor, grpc.aio.StreamUnaryClientInterceptor +): + """Async gRPC interceptor for stream-unary calls.""" + + async def intercept_stream_unary( + self, + continuation: Callable, + client_call_details: grpc.aio.ClientCallDetails, + request_iterator: Any, + ) -> Any: + return await continuation( + await self._add_metadata(client_call_details), request_iterator + ) + + +class _AsyncStreamStreamAuthInterceptor( + _AsyncBaseAuthInterceptor, grpc.aio.StreamStreamClientInterceptor +): + """Async gRPC interceptor for stream-stream calls.""" + + async def intercept_stream_stream( + self, + continuation: Callable, + client_call_details: grpc.aio.ClientCallDetails, + request_iterator: Any, + ) -> Any: + return await continuation( + await self._add_metadata(client_call_details), request_iterator + ) + + +class SpannerOmniCredentials(google.auth.credentials.Credentials): + """Credentials for Spanner Omni authentication using the OPAQUE protocol. + + Args: + username (str): The username for login. + password (str | bytes): The password for login. + target (str): The endpoint / target address for Spanner Omni. + use_plain_text (bool): Whether to use an insecure (plaintext) connection. + ca_certificate (str, optional): Path to the root CA certificate file. + client_certificate (str, optional): Path to the client certificate file for mTLS. + client_key (str, optional): Path to the client private key file for mTLS. + ssl_credentials (grpc.ChannelCredentials, optional): Pre-constructed SSL channel credentials. + """ + + def __init__( + self, + username: str, + password: str | bytes, + target: str, + use_plain_text: bool = False, + ca_certificate: Optional[str] = None, + client_certificate: Optional[str] = None, + client_key: Optional[str] = None, + ssl_credentials: Optional[grpc.ChannelCredentials] = None, + ) -> None: + super().__init__() + if not username: + raise ValueError("username cannot be empty") + if not password: + raise ValueError("password cannot be empty") + if not target: + raise ValueError("target cannot be empty") + + self.username = username + self._password: bytes = ( + password.encode("utf-8") if isinstance(password, str) else bytes(password) + ) + + # Parse target scheme + if target.startswith("http://"): + self.target = target[7:] + self.use_plain_text = True + _LOGGER.warning("Using plaintext connection for Spanner Omni credentials.") + elif target.startswith("https://"): + self.target = target[8:] + self.use_plain_text = use_plain_text + else: + self.target = target + self.use_plain_text = use_plain_text + + self.ca_certificate = ca_certificate + self.client_certificate = client_certificate + self.client_key = client_key + self.ssl_credentials = ssl_credentials + + self.token: Optional[str] = None + self.expiry: Optional[datetime.datetime] = None + self._lock = threading.Lock() + + def init_channel( + self, + use_plain_text: bool = False, + ca_certificate: Optional[str] = None, + client_certificate: Optional[str] = None, + client_key: Optional[str] = None, + ssl_credentials: Optional[grpc.ChannelCredentials] = None, + ) -> None: + """Initializes or updates channel TLS/transport settings.""" + self.use_plain_text = use_plain_text + if self.use_plain_text: + _LOGGER.warning("Using plaintext connection for Spanner Omni credentials.") + self.ca_certificate = ca_certificate + self.client_certificate = client_certificate + self.client_key = client_key + self.ssl_credentials = ssl_credentials + + def create_auth_interceptor(self, is_async: bool = False) -> Any: + """Creates a gRPC interceptor that attaches the Bearer token.""" + if is_async: + return self.create_async_auth_interceptors() + return _OmniAuthInterceptor(self) + + def create_async_auth_interceptors( + self, + ) -> Sequence[grpc.aio.ClientInterceptor]: + """Creates async gRPC interceptors that attach the Bearer token.""" + return [ + _AsyncUnaryUnaryAuthInterceptor(self), + _AsyncUnaryStreamAuthInterceptor(self), + _AsyncStreamUnaryAuthInterceptor(self), + _AsyncStreamStreamAuthInterceptor(self), + ] + + def create_async_auth_interceptor( + self, + ) -> Sequence[grpc.aio.ClientInterceptor]: + """Creates async gRPC interceptors that attach the Bearer token.""" + return self.create_async_auth_interceptors() + + def _perform_refresh_token(self, request: Any = None) -> None: + """Refreshes the access token by performing the OPAQUE login flow with Spanner Omni. + + Args: + request (Any, optional): Unused; part of google.auth.credentials.Credentials interface. + """ + with self._lock: + if self.valid: + return + login_channel = None + try: + if self.use_plain_text: + login_channel = grpc.insecure_channel(self.target) + elif self.ssl_credentials is not None: + login_channel = grpc.secure_channel( + self.target, self.ssl_credentials + ) + elif self.ca_certificate: + with open(self.ca_certificate, "rb") as f: + ca_cert = f.read() + if self.client_certificate and self.client_key: + with open(self.client_certificate, "rb") as f: + client_cert = f.read() + with open(self.client_key, "rb") as f: + private_key = f.read() + ssl_creds = grpc.ssl_channel_credentials( + root_certificates=ca_cert, + private_key=private_key, + certificate_chain=client_cert, + ) + elif self.client_certificate or self.client_key: + raise ValueError( + "Both client_certificate and client_key must be provided for mTLS" + ) + else: + ssl_creds = grpc.ssl_channel_credentials( + root_certificates=ca_cert + ) + login_channel = grpc.secure_channel(self.target, ssl_creds) + else: + raise ValueError( + "TLS/mTLS connection requires ca_certificate to be set for Spanner Omni" + ) + + client = LoginClient(login_channel) + proto_token = client.login(self.username, self._password) + + token_bytes = proto_token.SerializeToString() + self.token = base64.b64encode(token_bytes).decode("ascii") + + if proto_token.HasField("expiration_time") and ( + proto_token.expiration_time.seconds > 0 + or proto_token.expiration_time.nanos > 0 + ): + seconds = proto_token.expiration_time.seconds + nanos = proto_token.expiration_time.nanos + self.expiry = datetime.datetime.fromtimestamp( + seconds + nanos / 1e9, tz=datetime.timezone.utc + ).replace(tzinfo=None) + else: + self.expiry = datetime.datetime.now(datetime.timezone.utc).replace( + tzinfo=None + ) + datetime.timedelta(hours=1) + except (grpc.RpcError, google.auth.exceptions.RefreshError): + raise + except Exception as e: + raise google.auth.exceptions.RefreshError( + f"Failed to login to Spanner Omni: {e}" + ) from e + finally: + if login_channel is not None: + login_channel.close() + + def refresh(self, request: Any = None) -> None: + """Refreshes the access token. + + Args: + request (Any, optional): Unused; part of google.auth.credentials.Credentials interface. + """ + self._perform_refresh_token(request) + + def apply( + self, headers: MutableMapping[str, str], token: Optional[str] = None + ) -> None: + """Applies the access token to request headers.""" + headers["authorization"] = f"Bearer {token or self.token}" + + def before_request( + self, + request: Any, + method: str, + url: str, + headers: MutableMapping[str, str], + ) -> None: + """Performs token refresh if expired/missing and applies authorization header.""" + if not self.valid: + self.refresh(request) + self.apply(headers) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/login_client.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/login_client.py new file mode 100644 index 000000000000..398e0d877e25 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/login_client.py @@ -0,0 +1,150 @@ +# -*- 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. +# + +"""Client for Spanner Omni LoginService gRPC API.""" + +from __future__ import annotations + +import queue +from typing import Iterator, Optional + +import grpc + +from google.cloud.spanner_v1.omni.opaque import ( + EXPECTED_ENVELOPE_SIZE, + UserAuthenticator, +) +from google.cloud.spanner_v1.omni.proto import ( + authentication_pb2, + login_pb2, + login_pb2_grpc, +) + + +class _RequestIterator(Iterator[login_pb2.LoginRequest]): + """Thread-safe request iterator for gRPC bidirectional streaming.""" + + def __init__(self) -> None: + self._queue: queue.Queue[Optional[login_pb2.LoginRequest]] = queue.Queue() + self._closed = False + + def send(self, request: login_pb2.LoginRequest) -> None: + self._queue.put(request) + + def close(self) -> None: + if not self._closed: + self._closed = True + self._queue.put(None) + + def __iter__(self) -> _RequestIterator: + return self + + def __next__(self) -> login_pb2.LoginRequest: + item = self._queue.get() + if item is None: + raise StopIteration + return item + + +class LoginClient: + """Client for Spanner Omni LoginService.""" + + EXPECTED_ENVELOPE_SIZE = EXPECTED_ENVELOPE_SIZE + + def __init__(self, channel: grpc.Channel) -> None: + self._stub = login_pb2_grpc.LoginServiceStub(channel) + + def login( + self, username: str, password: str | bytes, timeout: float = 60.0 + ) -> login_pb2.AccessToken: + """Performs the full OPAQUE authentication handshake and returns an AccessToken. + + Args: + username (str): The username for login. + password (str | bytes): The password for login. + timeout (float): RPC timeout in seconds. + + Returns: + login_pb2.AccessToken: The issued access token proto. + + Raises: + ValueError: If handshake validation fails. + grpc.RpcError: If gRPC communication fails. + """ + if not username: + raise ValueError("username cannot be empty") + if not password: + raise ValueError("password cannot be empty") + + req_iterator = _RequestIterator() + try: + call = self._stub.Login(req_iterator, timeout=timeout) + + def _safe_next(stage: str) -> login_pb2.LoginResponse: + try: + return next(call) + except StopIteration: + raise ValueError(f"Server closed stream prematurely during {stage}") + + # Step 1: Handshake Request + handshake_req = login_pb2.LoginRequest( + username=username, + handshake_request=authentication_pb2.PasswordAuthenticationHandshakeRequest(), + ) + req_iterator.send(handshake_req) + handshake_resp = _safe_next("handshake") + + if not handshake_resp.HasField("handshake_response"): + raise ValueError("Failed to receive handshake response from server") + + method = handshake_resp.handshake_response.password_authentication_protocol + if ( + method + != authentication_pb2.PasswordAuthenticationProtocol.PASSWORD_AUTHENTICATION_PROTOCOL_OPAQUE + ): + raise ValueError( + f"Unsupported password authentication protocol: {method}" + ) + + if not handshake_resp.handshake_response.HasField("hash_parameters"): + raise ValueError("Handshake response missing hash_parameters") + + hash_params = handshake_resp.handshake_response.hash_parameters + authenticator = UserAuthenticator(username, password, hash_params) + + # Step 2: Initial OPAQUE Request + initial_req = authenticator.initial_request() + req_iterator.send(initial_req) + initial_resp = _safe_next("initial OPAQUE exchange") + + # Step 3: Final OPAQUE Request + final_req = authenticator.final_request(initial_resp) + req_iterator.send(final_req) + req_iterator.close() + + # Final Response with AccessToken + final_resp = _safe_next("final OPAQUE exchange") + if not final_resp.HasField("access_token"): + raise ValueError( + "Server failed to return an access token in final response" + ) + + return final_resp.access_token + except Exception: + if "call" in locals() and hasattr(call, "cancel"): + call.cancel() + req_iterator.close() + raise diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/opaque.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/opaque.py new file mode 100644 index 000000000000..27d261fa4a87 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/opaque.py @@ -0,0 +1,576 @@ +# -*- 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. +# + +"""OPAQUE protocol cryptographic utilities for Spanner Omni authentication.""" + +from __future__ import annotations + +import hashlib +import hmac +import secrets +from typing import Optional, Tuple + +from cryptography.hazmat.primitives import hashes +from cryptography.hazmat.primitives.kdf.argon2 import Argon2id +from cryptography.hazmat.primitives.kdf.hkdf import HKDF + +from google.cloud.spanner_v1.omni.proto import login_pb2 + +LOGIN_DOMAIN_SEPARATION_TAG = b"Spanner-Omni-Login" +AUTH_KEY_INFO = b"AuthKey" +EXPORT_KEY_INFO = b"ExportKey" +PRIVATE_KEY_INFO = b"PrivateKey" +MASKING_KEY_INFO = b"MaskingKey" +DIFFIE_HELLMAN_KEY_INFO = b"OPAQUE-DeriveDiffieHellmanKeyPair" + +NONCE_LENGTH = 32 +MAC_TAG_LENGTH = 32 +PUBLIC_KEY_LENGTH = 33 +EXPECTED_ENVELOPE_SIZE = PUBLIC_KEY_LENGTH + NONCE_LENGTH + MAC_TAG_LENGTH # 97 + +# NIST P-256 (secp256r1) Curve Constants +P = 0xFFFFFFFF00000001000000000000000000000000FFFFFFFFFFFFFFFFFFFFFFFF +A = (P - 3) % P +B = 0x5AC635D8AA3A93E7B3EBBD55769886BC651D06B0CC53B0F63BCE3C3E27D2604B +Z = (P - 10) % P +ORDER = 0xFFFFFFFF00000000FFFFFFFFFFFFFFFFBCE6FAADA7179E84F3B9CAC2FC632551 +GX = 0x6B17D1F2E12C4247F8BCE6E563A440F277037D812DEB33A0F4A13945D898C296 +GY = 0x4FE342E2FE1A7F9B8EE7EB4A7C0F9E162BCE33576B315ECECBB6406837BF51F5 +G = (GX, GY) +P_MINUS_1_OVER_2 = (P - 1) // 2 +P_PLUS_1_OVER_4 = (P + 1) // 4 + + +def _clear(b: Optional[bytearray]) -> None: + """Zeroizes a mutable bytearray in place.""" + if b is not None: + for i in range(len(b)): + b[i] = 0 + + +def point_add( + p1: Optional[Tuple[int, int]], p2: Optional[Tuple[int, int]] +) -> Optional[Tuple[int, int]]: + """Adds two points on the P-256 elliptic curve.""" + if p1 is None: + return p2 + if p2 is None: + return p1 + x1, y1 = p1 + x2, y2 = p2 + if x1 == x2: + if (y1 + y2) % P == 0: + return None + lam = ((3 * x1 * x1 + A) * pow(2 * y1, P - 2, P)) % P + else: + lam = ((y2 - y1) * pow(x2 - x1, P - 2, P)) % P + x3 = (lam * lam - x1 - x2) % P + y3 = (lam * (x1 - x3) - y1) % P + return (x3, y3) + + +def point_mul(pt: Optional[Tuple[int, int]], k: int) -> Optional[Tuple[int, int]]: + """Multiplies a point on the P-256 elliptic curve by a scalar k.""" + if pt is None or k == 0: + return None + res = None + curr = pt + while k > 0: + if k & 1: + res = point_add(res, curr) + curr = point_add(curr, curr) + k >>= 1 + return res + + +def marshal_compressed(pt: Optional[Tuple[int, int]]) -> bytes: + """Encodes a P-256 curve point into 33-byte compressed SEC1 format.""" + if pt is None: + raise ValueError("Point at infinity cannot be compressed") + x, y = pt + prefix = b"\x02" if (y % 2 == 0) else b"\x03" + return prefix + x.to_bytes(32, "big") + + +def unmarshal_compressed(data: bytes) -> Tuple[int, int]: + """Decodes a 33-byte compressed SEC1 format point on P-256.""" + if len(data) != 33: + raise ValueError(f"Invalid compressed point length: {len(data)}") + prefix = data[0] + if prefix not in (2, 3): + raise ValueError(f"Invalid compressed point prefix: {prefix}") + x = int.from_bytes(data[1:], "big") + if x >= P: + raise ValueError("x coordinate exceeds field prime") + rhs = (pow(x, 3, P) + A * x + B) % P + y = pow(rhs, P_PLUS_1_OVER_4, P) + if (pow(y, 2, P) - rhs) % P != 0: + raise ValueError("Point is not on curve") + if (y % 2) != (prefix & 1): + y = (P - y) % P + return (x, y) + + +def expand_message_xmd(msg: bytes, dst: bytes, len_in_bytes: int) -> bytes: + """Implements expand_message_xmd for SHA-256 per RFC 9380 Section 5.3.1.""" + if len(dst) > 255: + dst = hashlib.sha256(b"H2C-OVERSIZE-DST-" + dst).digest() + dst_len = bytes([len(dst)]) + b_in_bytes = 32 + ell = (len_in_bytes + b_in_bytes - 1) // b_in_bytes + z_pad = b"\x00" * 64 + lib_str = len_in_bytes.to_bytes(2, "big") + b0 = hashlib.sha256(z_pad + msg + lib_str + b"\x00" + dst + dst_len).digest() + b1 = hashlib.sha256(b0 + b"\x01" + dst + dst_len).digest() + res = bytearray(b1) + prev = b1 + for i in range(2, ell + 1): + tmp = bytes(x ^ y for x, y in zip(b0, prev)) + bi = hashlib.sha256(tmp + bytes([i]) + dst + dst_len).digest() + res.extend(bi) + prev = bi + return bytes(res[:len_in_bytes]) + + +def map_to_curve_sswu(u: int) -> Tuple[int, int]: + """Implements Simplified SWU mapping for P-256 per RFC 9380 Section 6.6.2.""" + u2 = (u * u) % P + tv1 = (Z * u2) % P + tv2 = (tv1 * tv1 + tv1) % P + tv3 = (B * (tv2 + 1)) % P + + if tv2 != 0: + tv4 = (-tv2) % P + else: + tv4 = Z + tv4 = (tv4 * A) % P + + tv4_inv = pow(tv4, P - 2, P) + x1 = (tv3 * tv4_inv) % P + + gx1 = (pow(x1, 3, P) + A * x1 + B) % P + e1 = pow(gx1, P_MINUS_1_OVER_2, P) + is_square = e1 == 1 or gx1 == 0 + + if is_square: + x = x1 + y = pow(gx1, P_PLUS_1_OVER_4, P) + else: + x = (tv1 * x1) % P + gx2 = (pow(x, 3, P) + A * x + B) % P + y = pow(gx2, P_PLUS_1_OVER_4, P) + + if (u & 1) != (y & 1): + y = (P - y) % P + return (x, y) + + +def hash_to_curve_p256(msg: bytes, dst: bytes) -> Tuple[int, int]: + """Implements P256_XMD:SHA-256_SSWU_RO_ hash-to-curve per RFC 9380 Section 8.2.""" + ub = expand_message_xmd(msg, dst, 96) + u0 = int.from_bytes(ub[:48], "big") % P + u1 = int.from_bytes(ub[48:], "big") % P + q0 = map_to_curve_sswu(u0) + q1 = map_to_curve_sswu(u1) + res = point_add(q0, q1) + if res is None: + raise ValueError("Hash to curve produced point at infinity") + return res + + +def nonce() -> bytes: + """Generates a 32-byte cryptographically secure random nonce.""" + return secrets.token_bytes(NONCE_LENGTH) + + +def sha256_hash(data: bytes) -> bytes: + """Computes the SHA-256 digest of input data.""" + return hashlib.sha256(data).digest() + + +def hmac_sha256(key: bytes, message: bytes) -> bytes: + """Computes HMAC-SHA-256.""" + return hmac.new(key, message, hashlib.sha256).digest() + + +def mac(key: bytes, data: bytes) -> bytes: + """Computes a 32-byte MAC tag using HMAC-SHA-256.""" + return hmac_sha256(key, data)[:MAC_TAG_LENGTH] + + +def xor_bytes(a: bytes, b: bytes) -> bytes: + """Computes bitwise XOR of two equal-length byte sequences.""" + if len(a) != len(b): + raise ValueError(f"Byte sequences must have equal length: {len(a)} != {len(b)}") + return (int.from_bytes(a, "big") ^ int.from_bytes(b, "big")).to_bytes(len(a), "big") + + +def concat(*arrays: bytes) -> bytes: + """Concatenates multiple byte sequences.""" + return b"".join(arrays) + + +def expand(input_key_material: bytes, info: bytes, size: int) -> bytes: + """Expands key material using HKDF with SHA-256.""" + hkdf = HKDF( + algorithm=hashes.SHA256(), + length=size, + salt=b"", + info=info, + ) + return hkdf.derive(input_key_material) + + +def extract(input_key_material: bytes) -> bytes: + """Extracts key material using HKDF with label 'Extract'.""" + return expand(input_key_material, b"Extract", 32) + + +def _validate_hash_parameters(hash_parameters) -> None: + if hash_parameters is None: + raise ValueError("hash_parameters cannot be None") + if hasattr(hash_parameters, "HasField"): + if not hash_parameters.HasField("argon2_id_parameters"): + raise ValueError( + "hash_parameters must contain non-nil argon2_id_parameters" + ) + argon2_params = hash_parameters.argon2_id_parameters + else: + argon2_params = getattr(hash_parameters, "argon2_id_parameters", None) + if argon2_params is None: + raise ValueError( + "hash_parameters must contain non-nil argon2_id_parameters" + ) + + if not (1 <= argon2_params.iteration_count <= 10): + raise ValueError( + f"Invalid Argon2Id iteration count: {argon2_params.iteration_count} (must be between 1 and 10)" + ) + if not (8 <= argon2_params.memory_usage <= 65536): + raise ValueError( + f"Invalid Argon2Id memory usage: {argon2_params.memory_usage} (must be between 8 and 65536 KB)" + ) + if not (1 <= argon2_params.parallelism <= 255): + raise ValueError( + f"Invalid Argon2Id parallelism: {argon2_params.parallelism} (must be between 1 and 255)" + ) + if not (1 <= argon2_params.hash_size <= 512): + raise ValueError( + f"Invalid Argon2Id hash size: {argon2_params.hash_size} (must be between 1 and 512)" + ) + + +def stretch(input_bytes: bytes, hash_parameters) -> bytes: + """Stretches the OPRF evaluation using Argon2id with server hash parameters.""" + _validate_hash_parameters(hash_parameters) + argon2_params = hash_parameters.argon2_id_parameters + + salt = expand(input_bytes, b"Stretch", int(argon2_params.hash_size)) + argon2 = Argon2id( + salt=salt, + length=int(argon2_params.hash_size), + iterations=int(argon2_params.iteration_count), + lanes=int(argon2_params.parallelism), + memory_cost=int(argon2_params.memory_usage), + ) + return argon2.derive(input_bytes) + + +def random_oracle_sha256(x: bytes, max_val: int) -> bytes: + """Iterative SHA-256 random oracle reduction mod max_val.""" + hash_output_length = 256 + output_bit_length = max_val.bit_length() + hash_output_length + iter_count = (output_bit_length + hash_output_length - 1) // hash_output_length + if iter_count > 255: + raise ValueError( + f"Domain bit length must not be greater than 65280: {output_bit_length}" + ) + excess_bit_count = (iter_count * hash_output_length) - output_bit_length + hash_output = 0 + for i in range(1, iter_count + 1): + hash_output <<= hash_output_length + bignum_bytes = bytes([i]) + x + hashed_string = hashlib.sha256(bignum_bytes).digest() + new_big_num = int.from_bytes(hashed_string, "big") + hash_output += new_big_num + + hash_output >>= excess_bit_count + hash_output %= max_val + + scalar_len = hash_output_length // 8 + max_len = (max_val.bit_length() + 7) // 8 + if max_len > scalar_len: + scalar_len = max_len + return hash_output.to_bytes(scalar_len, "big") + + +def derive_key_pair(seed: bytes, info: bytes) -> Tuple[bytes, bytes]: + """Derives an ECDH public/private keypair from a seed and info string.""" + derive_input = seed + info + priv_bytes = random_oracle_sha256(derive_input, ORDER) + priv_int = int.from_bytes(priv_bytes, "big") + if priv_int == 0: + priv_int = 1 + priv_bytes = (1).to_bytes(32, "big") + pub_point = point_mul(G, priv_int) + pub_bytes = marshal_compressed(pub_point) + return pub_bytes, priv_bytes + + +def diffie_hellman(priv_bytes: bytes, pub_bytes: bytes) -> bytes: + """Computes the ECDH shared point peer_pub * priv.""" + pt = unmarshal_compressed(pub_bytes) + priv_int = int.from_bytes(priv_bytes, "big") + shared_pt = point_mul(pt, priv_int) + return marshal_compressed(shared_pt) + + +def blind(password: bytes, blind_scalar: Optional[bytes] = None) -> Tuple[bytes, bytes]: + """Blinds the password point using a random or provided scalar.""" + if len(password) == 0: + raise ValueError("Password cannot be empty") + pt = hash_to_curve_p256(password, LOGIN_DOMAIN_SEPARATION_TAG) + if blind_scalar is None: + while True: + r = secrets.randbelow(ORDER) + if r != 0: + break + blind_scalar = r.to_bytes(32, "big") + else: + r = int.from_bytes(blind_scalar, "big") + blinded_pt = point_mul(pt, r) + return marshal_compressed(blinded_pt), blind_scalar + + +def finalize(blind_scalar: bytes, evaluated_message: bytes) -> bytes: + """Finalizes the OPRF output by multiplying with r^-1 mod Order.""" + if len(blind_scalar) == 0: + raise ValueError("Blind scalar cannot be empty") + r = int.from_bytes(blind_scalar, "big") + r_inv = pow(r, ORDER - 2, ORDER) + eval_pt = unmarshal_compressed(evaluated_message) + res_pt = point_mul(eval_pt, r_inv) + return marshal_compressed(res_pt) + + +def derive_secret( + input_key_material: bytes, label: bytes, transcript_hash: bytes +) -> bytes: + """Derives a secret labeled with 'OPAQUE-'.""" + info = b"OPAQUE-" + label + transcript_hash + return expand(input_key_material, info, 32) + + +def derive_shared_keys( + input_key_material: bytes, preamble: bytes +) -> Tuple[bytes, bytes, bytes]: + """Derives the km2 (server MAC), km3 (client MAC), and sessionKey.""" + prk = extract(input_key_material) + preamble_hash = sha256_hash(preamble) + handshake_secret = derive_secret(prk, b"HandshakeSecret", preamble_hash) + session_key = derive_secret(prk, b"SessionKey", preamble_hash) + km2 = derive_secret(handshake_secret, b"ServerMAC", b"") + km3 = derive_secret(handshake_secret, b"ClientMAC", b"") + return km2, km3, session_key + + +def recover_client( + username: str, + randomized_password: bytes, + envelope_nonce: bytes, + auth_tag: bytes, + server_public_key: bytes, +) -> Tuple[bytes, bytes]: + """Recovers the client's export key and private key from the envelope.""" + auth_key = expand(randomized_password, envelope_nonce + AUTH_KEY_INFO, 32) + export_key = expand(randomized_password, envelope_nonce + EXPORT_KEY_INFO, 32) + seed = expand(randomized_password, envelope_nonce + PRIVATE_KEY_INFO, 32) + _, client_private_key = derive_key_pair(seed, DIFFIE_HELLMAN_KEY_INFO) + + expected_tag = mac( + auth_key, envelope_nonce + server_public_key + username.encode("utf-8") + ) + if len(auth_tag) != len(expected_tag) or not hmac.compare_digest( + expected_tag, auth_tag + ): + raise ValueError("Auth tag mismatch") + return export_key, client_private_key + + +class UserAuthenticator: + """Manages the client state and key exchanges for OPAQUE login authentication.""" + + def __init__(self, username: str, password: str | bytes, hash_parameters): + if not username: + raise ValueError("username cannot be empty") + if isinstance(password, str): + password = password.encode("utf-8") + if len(password) == 0: + raise ValueError("password cannot be empty") + if hash_parameters is None: + raise ValueError("hash_parameters cannot be None") + _validate_hash_parameters(hash_parameters) + + self.username = username + self._password: Optional[bytearray] = bytearray(password) + self.hash_parameters = hash_parameters + + self._blind: Optional[bytearray] = None + self._client_nonce: Optional[bytes] = None + self._client_public_keyshare: Optional[bytes] = None + self._client_private_keyshare: Optional[bytearray] = None + + def initial_request(self) -> login_pb2.LoginRequest: + """Generates the initial OPAQUE login request.""" + if self._password is None: + raise ValueError("Authenticator already used or password not available") + + try: + blinded_message, blind_scalar = blind(self._password) + self._blind = bytearray(blind_scalar) + + self._client_nonce = nonce() + random_nonce = nonce() + + pub_key, priv_key = derive_key_pair(random_nonce, DIFFIE_HELLMAN_KEY_INFO) + self._client_public_keyshare = pub_key + self._client_private_keyshare = bytearray(priv_key) + + initial_opaque_req = login_pb2.InitialOpaqueLoginRequest( + blinded_message=blinded_message, + client_nonce=self._client_nonce, + client_public_keyshare=self._client_public_keyshare, + ) + opaque_req = login_pb2.OpaqueLoginRequest( + initial_request=initial_opaque_req + ) + return login_pb2.LoginRequest( + username=self.username, + opaque_request=opaque_req, + ) + finally: + _clear(self._password) + self._password = None + + def final_request( + self, initial_response: login_pb2.LoginResponse + ) -> login_pb2.LoginRequest: + """Generates the final OPAQUE login request containing the client MAC.""" + if initial_response is None: + raise ValueError("initial_response cannot be None") + if ( + self._client_public_keyshare is None + or self._client_nonce is None + or self._client_private_keyshare is None + or self._blind is None + ): + raise ValueError( + "Authenticator not initialized; initial_request must be called first" + ) + + if ( + not hasattr(initial_response, "HasField") + or not initial_response.HasField("opaque_response") + or not initial_response.opaque_response.HasField("initial_response") + ): + raise ValueError("Expected initial opaque response from server") + + initial_opaque_resp = initial_response.opaque_response.initial_response + + evaluated_message = initial_opaque_resp.evaluated_message + masking_nonce = initial_opaque_resp.masking_nonce + masked_response = initial_opaque_resp.masked_response + server_nonce = initial_opaque_resp.server_nonce + server_mac = initial_opaque_resp.server_mac + server_public_keyshare = initial_opaque_resp.server_public_keyshare + + if len(masked_response) != EXPECTED_ENVELOPE_SIZE: + raise ValueError( + f"Invalid masked response length: got {len(masked_response)}, want {EXPECTED_ENVELOPE_SIZE}" + ) + + try: + oprf = finalize(self._blind, evaluated_message) + stretched_oprf = stretch(oprf, self.hash_parameters) + randomized_password = extract(concat(oprf, stretched_oprf)) + + masking_key = expand(randomized_password, MASKING_KEY_INFO, 32) + credential_response_pad = expand( + masking_key, + concat(masking_nonce, b"CredentialResponsePad"), + len(masked_response), + ) + serialized_envelope = xor_bytes(masked_response, credential_response_pad) + if len(serialized_envelope) != EXPECTED_ENVELOPE_SIZE: + raise ValueError( + f"Invalid serialized envelope length: got {len(serialized_envelope)}, want {EXPECTED_ENVELOPE_SIZE}" + ) + + server_public_key = serialized_envelope[:PUBLIC_KEY_LENGTH] + envelope_nonce = serialized_envelope[ + PUBLIC_KEY_LENGTH : PUBLIC_KEY_LENGTH + NONCE_LENGTH + ] + auth_tag = serialized_envelope[ + PUBLIC_KEY_LENGTH + NONCE_LENGTH : EXPECTED_ENVELOPE_SIZE + ] + + export_key, client_private_key = recover_client( + self.username, + randomized_password, + envelope_nonce, + auth_tag, + server_public_key, + ) + + dh1 = diffie_hellman(self._client_private_keyshare, server_public_keyshare) + dh2 = diffie_hellman(self._client_private_keyshare, server_public_key) + dh3 = diffie_hellman(client_private_key, server_public_keyshare) + + input_key_material = concat(dh1, dh2, dh3) + + preamble = concat( + b"OPAQUEv1-", + self.username.encode("utf-8"), + self._client_nonce, + self._client_public_keyshare, + server_public_key, + evaluated_message, + server_nonce, + server_public_keyshare, + ) + + km2, km3, _ = derive_shared_keys(input_key_material, preamble) + + hashed_preamble = sha256_hash(preamble) + expected_server_mac = mac(km2, hashed_preamble) + if len(server_mac) != len(expected_server_mac) or not hmac.compare_digest( + expected_server_mac, server_mac + ): + raise ValueError("Server MAC mismatch") + + client_mac = mac(km3, sha256_hash(concat(preamble, expected_server_mac))) + + final_opaque_req = login_pb2.FinalOpaqueLoginRequest(client_mac=client_mac) + opaque_req = login_pb2.OpaqueLoginRequest(final_request=final_opaque_req) + return login_pb2.LoginRequest( + username=self.username, + opaque_request=opaque_req, + ) + finally: + _clear(self._blind) + self._blind = None + _clear(self._client_private_keyshare) + self._client_private_keyshare = None diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/__init__.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/__init__.py new file mode 100644 index 000000000000..b433d320f3b3 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/__init__.py @@ -0,0 +1,27 @@ +# -*- 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. +# + +from google.cloud.spanner_v1.omni.proto import ( + authentication_pb2, + login_pb2, + login_pb2_grpc, +) + +__all__ = ( + "authentication_pb2", + "login_pb2", + "login_pb2_grpc", +) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2.py new file mode 100644 index 000000000000..d1fed45a60b0 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2.py @@ -0,0 +1,53 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: google/cloud/spanner_v1/omni/proto/authentication.proto +# Protobuf Python Version: 7.35.1 +"""Generated protocol buffer code.""" + +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import runtime_version as _runtime_version +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder + +_runtime_version.ValidateProtobufRuntimeVersion( + _runtime_version.Domain.PUBLIC, + 6, + 33, + 5, + "", + "google/cloud/spanner_v1/omni/proto/authentication.proto", +) +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( + b'\n7google/cloud/spanner_v1/omni/proto/authentication.proto\x12\x16google.spanner.auth.v1"\xe6\x01\n\x0eHashParameters\x12Y\n\x14\x61rgon2_id_parameters\x18\x01 \x01(\x0b\x32\x39.google.spanner.auth.v1.HashParameters.Argon2IdParametersH\x00\x1ak\n\x12\x41rgon2IdParameters\x12\x17\n\x0fiteration_count\x18\x01 \x01(\r\x12\x14\n\x0cmemory_usage\x18\x02 \x01(\r\x12\x13\n\x0bparallelism\x18\x03 \x01(\r\x12\x11\n\thash_size\x18\x04 \x01(\rB\x0c\n\nparameters"(\n&PasswordAuthenticationHandshakeRequest"\xcc\x01\n\'PasswordAuthenticationHandshakeResponse\x12`\n password_authentication_protocol\x18\x01 \x01(\x0e\x32\x36.google.spanner.auth.v1.PasswordAuthenticationProtocol\x12?\n\x0fhash_parameters\x18\x02 \x01(\x0b\x32&.google.spanner.auth.v1.HashParameters*\x7f\n\x1ePasswordAuthenticationProtocol\x12\x30\n,PASSWORD_AUTHENTICATION_PROTOCOL_UNSPECIFIED\x10\x00\x12+\n\'PASSWORD_AUTHENTICATION_PROTOCOL_OPAQUE\x10\x02\x42\x43\n\x1d\x63om.google.cloud.spanner.omniP\x01Z cloud.google.com/go/spanner/omnib\x06proto3' +) + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages( + DESCRIPTOR, "google.cloud.spanner_v1.omni.proto.authentication_pb2", _globals +) +if not _descriptor._USE_C_DESCRIPTORS: + _globals["DESCRIPTOR"]._loaded_options = None + _globals[ + "DESCRIPTOR" + ]._serialized_options = ( + b"\n\035com.google.cloud.spanner.omniP\001Z cloud.google.com/go/spanner/omni" + ) + _globals["_PASSWORDAUTHENTICATIONPROTOCOL"]._serialized_start = 565 + _globals["_PASSWORDAUTHENTICATIONPROTOCOL"]._serialized_end = 692 + _globals["_HASHPARAMETERS"]._serialized_start = 84 + _globals["_HASHPARAMETERS"]._serialized_end = 314 + _globals["_HASHPARAMETERS_ARGON2IDPARAMETERS"]._serialized_start = 193 + _globals["_HASHPARAMETERS_ARGON2IDPARAMETERS"]._serialized_end = 300 + _globals["_PASSWORDAUTHENTICATIONHANDSHAKEREQUEST"]._serialized_start = 316 + _globals["_PASSWORDAUTHENTICATIONHANDSHAKEREQUEST"]._serialized_end = 356 + _globals["_PASSWORDAUTHENTICATIONHANDSHAKERESPONSE"]._serialized_start = 359 + _globals["_PASSWORDAUTHENTICATIONHANDSHAKERESPONSE"]._serialized_end = 563 +# @@protoc_insertion_point(module_scope) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2.pyi b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2.pyi new file mode 100644 index 000000000000..cb06f2bace00 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2.pyi @@ -0,0 +1,67 @@ +from collections.abc import Mapping as _Mapping +from typing import ClassVar as _ClassVar +from typing import Optional as _Optional +from typing import Union as _Union + +from google.protobuf import descriptor as _descriptor +from google.protobuf import message as _message +from google.protobuf.internal import enum_type_wrapper as _enum_type_wrapper + +DESCRIPTOR: _descriptor.FileDescriptor + +class PasswordAuthenticationProtocol(int, metaclass=_enum_type_wrapper.EnumTypeWrapper): + __slots__ = () + PASSWORD_AUTHENTICATION_PROTOCOL_UNSPECIFIED: _ClassVar[ + PasswordAuthenticationProtocol + ] + PASSWORD_AUTHENTICATION_PROTOCOL_OPAQUE: _ClassVar[PasswordAuthenticationProtocol] + +PASSWORD_AUTHENTICATION_PROTOCOL_UNSPECIFIED: PasswordAuthenticationProtocol +PASSWORD_AUTHENTICATION_PROTOCOL_OPAQUE: PasswordAuthenticationProtocol + +class HashParameters(_message.Message): + __slots__ = ("argon2_id_parameters",) + class Argon2IdParameters(_message.Message): + __slots__ = ("iteration_count", "memory_usage", "parallelism", "hash_size") + ITERATION_COUNT_FIELD_NUMBER: _ClassVar[int] + MEMORY_USAGE_FIELD_NUMBER: _ClassVar[int] + PARALLELISM_FIELD_NUMBER: _ClassVar[int] + HASH_SIZE_FIELD_NUMBER: _ClassVar[int] + iteration_count: int + memory_usage: int + parallelism: int + hash_size: int + def __init__( + self, + iteration_count: _Optional[int] = ..., + memory_usage: _Optional[int] = ..., + parallelism: _Optional[int] = ..., + hash_size: _Optional[int] = ..., + ) -> None: ... + + ARGON2_ID_PARAMETERS_FIELD_NUMBER: _ClassVar[int] + argon2_id_parameters: HashParameters.Argon2IdParameters + def __init__( + self, + argon2_id_parameters: _Optional[ + _Union[HashParameters.Argon2IdParameters, _Mapping] + ] = ..., + ) -> None: ... + +class PasswordAuthenticationHandshakeRequest(_message.Message): + __slots__ = () + def __init__(self) -> None: ... + +class PasswordAuthenticationHandshakeResponse(_message.Message): + __slots__ = ("password_authentication_protocol", "hash_parameters") + PASSWORD_AUTHENTICATION_PROTOCOL_FIELD_NUMBER: _ClassVar[int] + HASH_PARAMETERS_FIELD_NUMBER: _ClassVar[int] + password_authentication_protocol: PasswordAuthenticationProtocol + hash_parameters: HashParameters + def __init__( + self, + password_authentication_protocol: _Optional[ + _Union[PasswordAuthenticationProtocol, str] + ] = ..., + hash_parameters: _Optional[_Union[HashParameters, _Mapping]] = ..., + ) -> None: ... diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2_grpc.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2_grpc.py new file mode 100644 index 000000000000..ef4c49f28e09 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/authentication_pb2_grpc.py @@ -0,0 +1,26 @@ +# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! +"""Client and server classes corresponding to protobuf-defined services.""" + +import grpc + +GRPC_GENERATED_VERSION = "1.59.0" +GRPC_VERSION = grpc.__version__ +_version_not_supported = False + +try: + from grpc._utilities import first_version_is_lower + + _version_not_supported = first_version_is_lower( + GRPC_VERSION, GRPC_GENERATED_VERSION + ) +except ImportError: + _version_not_supported = False + +if _version_not_supported: + raise RuntimeError( + f"The grpc package installed is at version {GRPC_VERSION}," + + " but the generated code in google/cloud/spanner_v1/omni/proto/authentication_pb2_grpc.py depends on" + + f" grpcio>={GRPC_GENERATED_VERSION}." + + f" Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}" + + f" or downgrade your generated code using grpcio-tools<={GRPC_VERSION}." + ) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2.py new file mode 100644 index 000000000000..8df1d65700bb --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2.py @@ -0,0 +1,71 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: google/cloud/spanner_v1/omni/proto/login.proto +# Protobuf Python Version: 7.35.1 +"""Generated protocol buffer code.""" + +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import runtime_version as _runtime_version +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder + +_runtime_version.ValidateProtobufRuntimeVersion( + _runtime_version.Domain.PUBLIC, + 6, + 33, + 5, + "", + "google/cloud/spanner_v1/omni/proto/login.proto", +) +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + +from google.protobuf import timestamp_pb2 as google_dot_protobuf_dot_timestamp__pb2 + +from google.cloud.spanner_v1.omni.proto import ( + authentication_pb2 as google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_authentication__pb2, +) + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( + b'\n.google/cloud/spanner_v1/omni/proto/login.proto\x12\x16google.spanner.auth.v1\x1a\x1fgoogle/protobuf/timestamp.proto\x1a\x37google/cloud/spanner_v1/omni/proto/authentication.proto"\xe5\x02\n\x0b\x41\x63\x63\x65ssToken\x12\x10\n\x08username\x18\x01 \x01(\t\x12\x31\n\rcreation_time\x18\x02 \x01(\x0b\x32\x1a.google.protobuf.Timestamp\x12\x33\n\x0f\x65xpiration_time\x18\x03 \x01(\x0b\x32\x1a.google.protobuf.Timestamp\x12\x11\n\tsignature\x18\x04 \x01(\x0c\x12\x0e\n\x06key_id\x18\x05 \x01(\x03\x12N\n\x11\x61\x63\x63\x65ss_token_type\x18\x06 \x01(\x0e\x32\x33.google.spanner.auth.v1.AccessToken.AccessTokenType"i\n\x0f\x41\x63\x63\x65ssTokenType\x12!\n\x1d\x41\x43\x43\x45SS_TOKEN_TYPE_UNSPECIFIED\x10\x00\x12\x19\n\x15\x41\x43\x43\x45SS_TOKEN_TYPE_API\x10\x01\x12\x18\n\x14\x41\x43\x43\x45SS_TOKEN_TYPE_UI\x10\x02"j\n\x19InitialOpaqueLoginRequest\x12\x17\n\x0f\x62linded_message\x18\x01 \x01(\x0c\x12\x14\n\x0c\x63lient_nonce\x18\x02 \x01(\x0c\x12\x1e\n\x16\x63lient_public_keyshare\x18\x03 \x01(\x0c"-\n\x17\x46inalOpaqueLoginRequest\x12\x12\n\nclient_mac\x18\x01 \x01(\x0c"\xb1\x01\n\x1aInitialOpaqueLoginResponse\x12\x14\n\x0cserver_nonce\x18\x01 \x01(\x0c\x12\x1e\n\x16server_public_keyshare\x18\x02 \x01(\x0c\x12\x12\n\nserver_mac\x18\x03 \x01(\x0c\x12\x19\n\x11\x65valuated_message\x18\x04 \x01(\x0c\x12\x15\n\rmasking_nonce\x18\x05 \x01(\x0c\x12\x17\n\x0fmasked_response\x18\x06 \x01(\x0c"\xb7\x01\n\x12OpaqueLoginRequest\x12L\n\x0finitial_request\x18\x01 \x01(\x0b\x32\x31.google.spanner.auth.v1.InitialOpaqueLoginRequestH\x00\x12H\n\rfinal_request\x18\x02 \x01(\x0b\x32/.google.spanner.auth.v1.FinalOpaqueLoginRequestH\x00\x42\t\n\x07request"\xd7\x01\n\x13OpaqueLoginResponse\x12N\n\x10initial_response\x18\x01 \x01(\x0b\x32\x32.google.spanner.auth.v1.InitialOpaqueLoginResponseH\x00\x12S\n\x0e\x66inal_response\x18\x02 \x01(\x0b\x32\x39.google.spanner.auth.v1.OpaqueLoginResponse.FinalResponseH\x00\x1a\x0f\n\rFinalResponseB\n\n\x08response"\xce\x01\n\x0cLoginRequest\x12\x10\n\x08username\x18\x01 \x01(\t\x12\x44\n\x0eopaque_request\x18\x04 \x01(\x0b\x32*.google.spanner.auth.v1.OpaqueLoginRequestH\x00\x12[\n\x11handshake_request\x18\x05 \x01(\x0b\x32>.google.spanner.auth.v1.PasswordAuthenticationHandshakeRequestH\x00\x42\t\n\x07request"\xfd\x01\n\rLoginResponse\x12\x39\n\x0c\x61\x63\x63\x65ss_token\x18\x01 \x01(\x0b\x32#.google.spanner.auth.v1.AccessToken\x12\x46\n\x0fopaque_response\x18\x04 \x01(\x0b\x32+.google.spanner.auth.v1.OpaqueLoginResponseH\x00\x12]\n\x12handshake_response\x18\x05 \x01(\x0b\x32?.google.spanner.auth.v1.PasswordAuthenticationHandshakeResponseH\x00\x42\n\n\x08response2h\n\x0cLoginService\x12X\n\x05Login\x12$.google.spanner.auth.v1.LoginRequest\x1a%.google.spanner.auth.v1.LoginResponse(\x01\x30\x01\x42\x43\n\x1d\x63om.google.cloud.spanner.omniP\x01Z cloud.google.com/go/spanner/omnib\x06proto3' +) + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages( + DESCRIPTOR, "google.cloud.spanner_v1.omni.proto.login_pb2", _globals +) +if not _descriptor._USE_C_DESCRIPTORS: + _globals["DESCRIPTOR"]._loaded_options = None + _globals[ + "DESCRIPTOR" + ]._serialized_options = ( + b"\n\035com.google.cloud.spanner.omniP\001Z cloud.google.com/go/spanner/omni" + ) + _globals["_ACCESSTOKEN"]._serialized_start = 165 + _globals["_ACCESSTOKEN"]._serialized_end = 522 + _globals["_ACCESSTOKEN_ACCESSTOKENTYPE"]._serialized_start = 417 + _globals["_ACCESSTOKEN_ACCESSTOKENTYPE"]._serialized_end = 522 + _globals["_INITIALOPAQUELOGINREQUEST"]._serialized_start = 524 + _globals["_INITIALOPAQUELOGINREQUEST"]._serialized_end = 630 + _globals["_FINALOPAQUELOGINREQUEST"]._serialized_start = 632 + _globals["_FINALOPAQUELOGINREQUEST"]._serialized_end = 677 + _globals["_INITIALOPAQUELOGINRESPONSE"]._serialized_start = 680 + _globals["_INITIALOPAQUELOGINRESPONSE"]._serialized_end = 857 + _globals["_OPAQUELOGINREQUEST"]._serialized_start = 860 + _globals["_OPAQUELOGINREQUEST"]._serialized_end = 1043 + _globals["_OPAQUELOGINRESPONSE"]._serialized_start = 1046 + _globals["_OPAQUELOGINRESPONSE"]._serialized_end = 1261 + _globals["_OPAQUELOGINRESPONSE_FINALRESPONSE"]._serialized_start = 1234 + _globals["_OPAQUELOGINRESPONSE_FINALRESPONSE"]._serialized_end = 1249 + _globals["_LOGINREQUEST"]._serialized_start = 1264 + _globals["_LOGINREQUEST"]._serialized_end = 1470 + _globals["_LOGINRESPONSE"]._serialized_start = 1473 + _globals["_LOGINRESPONSE"]._serialized_end = 1726 + _globals["_LOGINSERVICE"]._serialized_start = 1728 + _globals["_LOGINSERVICE"]._serialized_end = 1832 +# @@protoc_insertion_point(module_scope) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2.pyi b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2.pyi new file mode 100644 index 000000000000..96bed67256b9 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2.pyi @@ -0,0 +1,176 @@ +import datetime +from collections.abc import Mapping as _Mapping +from typing import ClassVar as _ClassVar +from typing import Optional as _Optional +from typing import Union as _Union + +from google.protobuf import descriptor as _descriptor +from google.protobuf import message as _message +from google.protobuf import timestamp_pb2 as _timestamp_pb2 +from google.protobuf.internal import enum_type_wrapper as _enum_type_wrapper + +from google.cloud.spanner_v1.omni.proto import authentication_pb2 as _authentication_pb2 + +DESCRIPTOR: _descriptor.FileDescriptor + +class AccessToken(_message.Message): + __slots__ = ( + "username", + "creation_time", + "expiration_time", + "signature", + "key_id", + "access_token_type", + ) + class AccessTokenType(int, metaclass=_enum_type_wrapper.EnumTypeWrapper): + __slots__ = () + ACCESS_TOKEN_TYPE_UNSPECIFIED: _ClassVar[AccessToken.AccessTokenType] + ACCESS_TOKEN_TYPE_API: _ClassVar[AccessToken.AccessTokenType] + ACCESS_TOKEN_TYPE_UI: _ClassVar[AccessToken.AccessTokenType] + + ACCESS_TOKEN_TYPE_UNSPECIFIED: AccessToken.AccessTokenType + ACCESS_TOKEN_TYPE_API: AccessToken.AccessTokenType + ACCESS_TOKEN_TYPE_UI: AccessToken.AccessTokenType + USERNAME_FIELD_NUMBER: _ClassVar[int] + CREATION_TIME_FIELD_NUMBER: _ClassVar[int] + EXPIRATION_TIME_FIELD_NUMBER: _ClassVar[int] + SIGNATURE_FIELD_NUMBER: _ClassVar[int] + KEY_ID_FIELD_NUMBER: _ClassVar[int] + ACCESS_TOKEN_TYPE_FIELD_NUMBER: _ClassVar[int] + username: str + creation_time: _timestamp_pb2.Timestamp + expiration_time: _timestamp_pb2.Timestamp + signature: bytes + key_id: int + access_token_type: AccessToken.AccessTokenType + def __init__( + self, + username: _Optional[str] = ..., + creation_time: _Optional[ + _Union[datetime.datetime, _timestamp_pb2.Timestamp, _Mapping] + ] = ..., + expiration_time: _Optional[ + _Union[datetime.datetime, _timestamp_pb2.Timestamp, _Mapping] + ] = ..., + signature: _Optional[bytes] = ..., + key_id: _Optional[int] = ..., + access_token_type: _Optional[_Union[AccessToken.AccessTokenType, str]] = ..., + ) -> None: ... + +class InitialOpaqueLoginRequest(_message.Message): + __slots__ = ("blinded_message", "client_nonce", "client_public_keyshare") + BLINDED_MESSAGE_FIELD_NUMBER: _ClassVar[int] + CLIENT_NONCE_FIELD_NUMBER: _ClassVar[int] + CLIENT_PUBLIC_KEYSHARE_FIELD_NUMBER: _ClassVar[int] + blinded_message: bytes + client_nonce: bytes + client_public_keyshare: bytes + def __init__( + self, + blinded_message: _Optional[bytes] = ..., + client_nonce: _Optional[bytes] = ..., + client_public_keyshare: _Optional[bytes] = ..., + ) -> None: ... + +class FinalOpaqueLoginRequest(_message.Message): + __slots__ = ("client_mac",) + CLIENT_MAC_FIELD_NUMBER: _ClassVar[int] + client_mac: bytes + def __init__(self, client_mac: _Optional[bytes] = ...) -> None: ... + +class InitialOpaqueLoginResponse(_message.Message): + __slots__ = ( + "server_nonce", + "server_public_keyshare", + "server_mac", + "evaluated_message", + "masking_nonce", + "masked_response", + ) + SERVER_NONCE_FIELD_NUMBER: _ClassVar[int] + SERVER_PUBLIC_KEYSHARE_FIELD_NUMBER: _ClassVar[int] + SERVER_MAC_FIELD_NUMBER: _ClassVar[int] + EVALUATED_MESSAGE_FIELD_NUMBER: _ClassVar[int] + MASKING_NONCE_FIELD_NUMBER: _ClassVar[int] + MASKED_RESPONSE_FIELD_NUMBER: _ClassVar[int] + server_nonce: bytes + server_public_keyshare: bytes + server_mac: bytes + evaluated_message: bytes + masking_nonce: bytes + masked_response: bytes + def __init__( + self, + server_nonce: _Optional[bytes] = ..., + server_public_keyshare: _Optional[bytes] = ..., + server_mac: _Optional[bytes] = ..., + evaluated_message: _Optional[bytes] = ..., + masking_nonce: _Optional[bytes] = ..., + masked_response: _Optional[bytes] = ..., + ) -> None: ... + +class OpaqueLoginRequest(_message.Message): + __slots__ = ("initial_request", "final_request") + INITIAL_REQUEST_FIELD_NUMBER: _ClassVar[int] + FINAL_REQUEST_FIELD_NUMBER: _ClassVar[int] + initial_request: InitialOpaqueLoginRequest + final_request: FinalOpaqueLoginRequest + def __init__( + self, + initial_request: _Optional[_Union[InitialOpaqueLoginRequest, _Mapping]] = ..., + final_request: _Optional[_Union[FinalOpaqueLoginRequest, _Mapping]] = ..., + ) -> None: ... + +class OpaqueLoginResponse(_message.Message): + __slots__ = ("initial_response", "final_response") + class FinalResponse(_message.Message): + __slots__ = () + def __init__(self) -> None: ... + + INITIAL_RESPONSE_FIELD_NUMBER: _ClassVar[int] + FINAL_RESPONSE_FIELD_NUMBER: _ClassVar[int] + initial_response: InitialOpaqueLoginResponse + final_response: OpaqueLoginResponse.FinalResponse + def __init__( + self, + initial_response: _Optional[_Union[InitialOpaqueLoginResponse, _Mapping]] = ..., + final_response: _Optional[ + _Union[OpaqueLoginResponse.FinalResponse, _Mapping] + ] = ..., + ) -> None: ... + +class LoginRequest(_message.Message): + __slots__ = ("username", "opaque_request", "handshake_request") + USERNAME_FIELD_NUMBER: _ClassVar[int] + OPAQUE_REQUEST_FIELD_NUMBER: _ClassVar[int] + HANDSHAKE_REQUEST_FIELD_NUMBER: _ClassVar[int] + username: str + opaque_request: OpaqueLoginRequest + handshake_request: _authentication_pb2.PasswordAuthenticationHandshakeRequest + def __init__( + self, + username: _Optional[str] = ..., + opaque_request: _Optional[_Union[OpaqueLoginRequest, _Mapping]] = ..., + handshake_request: _Optional[ + _Union[_authentication_pb2.PasswordAuthenticationHandshakeRequest, _Mapping] + ] = ..., + ) -> None: ... + +class LoginResponse(_message.Message): + __slots__ = ("access_token", "opaque_response", "handshake_response") + ACCESS_TOKEN_FIELD_NUMBER: _ClassVar[int] + OPAQUE_RESPONSE_FIELD_NUMBER: _ClassVar[int] + HANDSHAKE_RESPONSE_FIELD_NUMBER: _ClassVar[int] + access_token: AccessToken + opaque_response: OpaqueLoginResponse + handshake_response: _authentication_pb2.PasswordAuthenticationHandshakeResponse + def __init__( + self, + access_token: _Optional[_Union[AccessToken, _Mapping]] = ..., + opaque_response: _Optional[_Union[OpaqueLoginResponse, _Mapping]] = ..., + handshake_response: _Optional[ + _Union[ + _authentication_pb2.PasswordAuthenticationHandshakeResponse, _Mapping + ] + ] = ..., + ) -> None: ... diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2_grpc.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2_grpc.py new file mode 100644 index 000000000000..18cc0685b2c2 --- /dev/null +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/omni/proto/login_pb2_grpc.py @@ -0,0 +1,109 @@ +# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! +"""Client and server classes corresponding to protobuf-defined services.""" + +import grpc + +from google.cloud.spanner_v1.omni.proto import ( + login_pb2 as google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_login__pb2, +) + +GRPC_GENERATED_VERSION = "1.59.0" +GRPC_VERSION = grpc.__version__ +_version_not_supported = False + +try: + from grpc._utilities import first_version_is_lower + + _version_not_supported = first_version_is_lower( + GRPC_VERSION, GRPC_GENERATED_VERSION + ) +except ImportError: + _version_not_supported = False + +if _version_not_supported: + raise RuntimeError( + f"The grpc package installed is at version {GRPC_VERSION}," + + " but the generated code in google/cloud/spanner_v1/omni/proto/login_pb2_grpc.py depends on" + + f" grpcio>={GRPC_GENERATED_VERSION}." + + f" Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}" + + f" or downgrade your generated code using grpcio-tools<={GRPC_VERSION}." + ) + + +class LoginServiceStub: + """The LoginService is used to authenticate users.""" + + def __init__(self, channel): + """Constructor. + + Args: + channel: A grpc.Channel. + """ + self.Login = channel.stream_stream( + "/google.spanner.auth.v1.LoginService/Login", + request_serializer=google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_login__pb2.LoginRequest.SerializeToString, + response_deserializer=google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_login__pb2.LoginResponse.FromString, + _registered_method=True, + ) + + +class LoginServiceServicer: + """The LoginService is used to authenticate users.""" + + def Login(self, request_iterator, context): + """Performs the login for Spanner Omni.""" + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details("Method not implemented!") + raise NotImplementedError("Method not implemented!") + + +def add_LoginServiceServicer_to_server(servicer, server): + rpc_method_handlers = { + "Login": grpc.stream_stream_rpc_method_handler( + servicer.Login, + request_deserializer=google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_login__pb2.LoginRequest.FromString, + response_serializer=google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_login__pb2.LoginResponse.SerializeToString, + ), + } + generic_handler = grpc.method_handlers_generic_handler( + "google.spanner.auth.v1.LoginService", rpc_method_handlers + ) + server.add_generic_rpc_handlers((generic_handler,)) + server.add_registered_method_handlers( + "google.spanner.auth.v1.LoginService", rpc_method_handlers + ) + + +# This class is part of an EXPERIMENTAL API. +class LoginService: + """The LoginService is used to authenticate users.""" + + @staticmethod + def Login( + request_iterator, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None, + ): + return grpc.experimental.stream_stream( + request_iterator, + target, + "/google.spanner.auth.v1.LoginService/Login", + google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_login__pb2.LoginRequest.SerializeToString, + google_dot_cloud_dot_spanner__v1_dot_omni_dot_proto_dot_login__pb2.LoginResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True, + ) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/testing/database_test.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/testing/database_test.py index 523946ab3545..1e11b7ee2230 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/testing/database_test.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/testing/database_test.py @@ -103,6 +103,7 @@ def spanner_api(self): client._client_certificate, client._client_key, self._interceptors, + credentials=client.credentials, ) self._spanner_api = SpannerClient( client_info=client_info, diff --git a/packages/google-cloud-spanner/setup.py b/packages/google-cloud-spanner/setup.py index b4e8a3908efc..c1a195e823f9 100644 --- a/packages/google-cloud-spanner/setup.py +++ b/packages/google-cloud-spanner/setup.py @@ -61,6 +61,7 @@ "opentelemetry-resourcedetector-gcp >= 1.8.0a0", "google-cloud-monitoring >= 2.28.0", "mmh3 >= 4.1.0", + "cryptography >= 44.0.0", ] extras = { "libcst": "libcst >= 0.2.5", diff --git a/packages/google-cloud-spanner/testing/constraints-3.10.txt b/packages/google-cloud-spanner/testing/constraints-3.10.txt index 2df6db04da12..97e2463b1398 100644 --- a/packages/google-cloud-spanner/testing/constraints-3.10.txt +++ b/packages/google-cloud-spanner/testing/constraints-3.10.txt @@ -22,3 +22,4 @@ google-cloud-monitoring==2.28.0 mmh3==4.1.0 libcst==0.2.5 googleapis-common-protos==1.69.2 +cryptography==44.0.0 diff --git a/packages/google-cloud-spanner/tests/system/_async/conftest.py b/packages/google-cloud-spanner/tests/system/_async/conftest.py index 0a09a9676074..b01db2579452 100644 --- a/packages/google-cloud-spanner/tests/system/_async/conftest.py +++ b/packages/google-cloud-spanner/tests/system/_async/conftest.py @@ -38,6 +38,8 @@ def spanner_client(): client_key=_helpers.CLIENT_KEY, client_options={"api_endpoint": _helpers.SPANNER_OMNI}, instance_type="omni", + username=_helpers.SPANNER_OMNI_USER, + password=_helpers.SPANNER_OMNI_PASSWORD, ) else: client_options = {"api_endpoint": _helpers.API_ENDPOINT} diff --git a/packages/google-cloud-spanner/tests/system/_helpers.py b/packages/google-cloud-spanner/tests/system/_helpers.py index 1bad9adec3da..57a1e47fb62b 100644 --- a/packages/google-cloud-spanner/tests/system/_helpers.py +++ b/packages/google-cloud-spanner/tests/system/_helpers.py @@ -68,6 +68,11 @@ CLIENT_KEY = os.getenv(CLIENT_KEY_ENVVAR) USE_PLAIN_TEXT = CA_CERTIFICATE is None +SPANNER_OMNI_USER_ENVVAR = "SPANNER_OMNI_USER" +SPANNER_OMNI_USER = os.getenv(SPANNER_OMNI_USER_ENVVAR) +SPANNER_OMNI_PASSWORD_ENVVAR = "SPANNER_OMNI_PASSWORD" +SPANNER_OMNI_PASSWORD = os.getenv(SPANNER_OMNI_PASSWORD_ENVVAR) + SPANNER_OMNI_INSTANCE = "default" DDL_STATEMENTS = ( diff --git a/packages/google-cloud-spanner/tests/system/conftest.py b/packages/google-cloud-spanner/tests/system/conftest.py index 5aaae4b4484d..757068bebaa5 100644 --- a/packages/google-cloud-spanner/tests/system/conftest.py +++ b/packages/google-cloud-spanner/tests/system/conftest.py @@ -124,6 +124,8 @@ def spanner_client(): client_key=_helpers.CLIENT_KEY, client_options={"api_endpoint": _helpers.SPANNER_OMNI}, instance_type="omni", + username=_helpers.SPANNER_OMNI_USER, + password=_helpers.SPANNER_OMNI_PASSWORD, ) else: client_options = {"api_endpoint": _helpers.API_ENDPOINT} diff --git a/packages/google-cloud-spanner/tests/system/test_dbapi.py b/packages/google-cloud-spanner/tests/system/test_dbapi.py index 82aada69e70a..53d9180663f9 100644 --- a/packages/google-cloud-spanner/tests/system/test_dbapi.py +++ b/packages/google-cloud-spanner/tests/system/test_dbapi.py @@ -1502,6 +1502,8 @@ def test_user_agent(self, shared_instance, dbapi_database): ca_certificate=_helpers.CA_CERTIFICATE, client_certificate=_helpers.CLIENT_CERTIFICATE, client_key=_helpers.CLIENT_KEY, + username=_helpers.SPANNER_OMNI_USER, + password=_helpers.SPANNER_OMNI_PASSWORD, ) assert ( conn.instance._client._client_info.user_agent diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_client.py b/packages/google-cloud-spanner/tests/unit/_async/test_client.py index 60bc98addc8e..9bf7a3751a87 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_client.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_client.py @@ -936,3 +936,123 @@ async def test_constructor_w_invalid_instance_type_raises_value_error(self): "instance_type must be one of 'cloud' or 'omni'", str(ctx.exception), ) + + @CrossSync.pytest + async def test_constructor_w_omni_username_password(self): + from google.cloud.spanner_v1._async.client import InstanceType + from google.cloud.spanner_v1.omni.credentials import ( + SpannerOmniCredentials, + ) + + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + password="test_password", + ) + self.assertEqual(client.project, "default") + self.assertEqual(client.instance_type, InstanceType.OMNI) + self.assertEqual(client._host, "omni-host:15000") + self.assertIsInstance(client._credentials, SpannerOmniCredentials) + self.assertEqual(client._credentials.username, "test_user") + self.assertEqual(client._credentials.target, "omni-host:15000") + + @CrossSync.pytest + async def test_constructor_w_omni_partial_credentials_raises_value_error(self): + from google.cloud.spanner_v1._async.client import InstanceType + + with self.assertRaises(ValueError) as ctx: + self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + ) + self.assertIn( + "Both username and password must be specified for Omni authentication", + str(ctx.exception), + ) + + with self.assertRaises(ValueError) as ctx: + self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + password="test_password", + ) + self.assertIn( + "Both username and password must be specified for Omni authentication", + str(ctx.exception), + ) + + @CrossSync.pytest + async def test_constructor_w_username_password_on_cloud_raises_value_error(self): + creds = build_scoped_credentials() + with self.assertRaises(ValueError) as ctx: + self._make_one( + project=self.PROJECT, + credentials=creds, + username="test_user", + password="test_password", + ) + self.assertIn( + "username and password can only be used when instance_type='omni'.", + str(ctx.exception), + ) + + @CrossSync.pytest + async def test_instance_admin_api_omni(self): + from google.cloud.spanner_v1._async.client import InstanceType + + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + password="test_password", + use_plain_text=True, + ) + + inst_module = "google.cloud.spanner_v1._async.client.InstanceAdminClient" + with mock.patch(inst_module) as instance_admin_client: + api = client.instance_admin_api + self.assertIs(api, instance_admin_client.return_value) + instance_admin_client.assert_called_once() + called_kw = instance_admin_client.call_args[1] + self.assertIn("transport", called_kw) + + @CrossSync.pytest + async def test_database_admin_api_omni(self): + from google.cloud.spanner_v1._async.client import InstanceType + + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + password="test_password", + use_plain_text=True, + ) + + db_module = "google.cloud.spanner_v1._async.client.DatabaseAdminClient" + with mock.patch(db_module) as database_admin_client: + api = client.database_admin_api + self.assertIs(api, database_admin_client.return_value) + database_admin_client.assert_called_once() + called_kw = database_admin_client.call_args[1] + self.assertIn("transport", called_kw) + + @CrossSync.pytest + async def test_constructor_w_omni_explicit_credentials_instance(self): + from google.cloud.spanner_v1._async.client import InstanceType + from google.cloud.spanner_v1.omni.credentials import SpannerOmniCredentials + + creds = SpannerOmniCredentials("user", "pass", "omni-host:15000") + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + credentials=creds, + ) + self.assertIs(client._credentials, creds) diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py b/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py index c49ada5ec9c6..883fe2840adc 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py @@ -141,3 +141,161 @@ async def test_create_experimental_host_transport_errors(self): MUT._create_experimental_host_transport( InstanceAdminGrpcTransport, "host", False, None, None, None ) + + async def test_create_spanner_omni_transport(self): + mock_factory = mock.MagicMock() + + # Plaintext with create_async_auth_interceptors + mock_creds1 = mock.MagicMock() + mock_creds1.create_async_auth_interceptors.return_value = ["interceptor1"] + with mock.patch("grpc.aio.insecure_channel") as mock_insecure: + MUT._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + credentials=mock_creds1, + ) + mock_insecure.assert_called_once_with( + target="localhost:9010", interceptors=["interceptor1"] + ) + + # Credentials with create_async_auth_interceptor returning list and single item + mock_creds2 = mock.MagicMock(spec=["create_async_auth_interceptor"]) + mock_creds2.create_async_auth_interceptor.return_value = ["interceptor2"] + with mock.patch("grpc.aio.insecure_channel") as mock_insecure: + MUT._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + credentials=mock_creds2, + ) + mock_insecure.assert_called_once_with( + target="localhost:9010", interceptors=["interceptor2"] + ) + + mock_creds2.create_async_auth_interceptor.return_value = "single_interceptor" + with mock.patch("grpc.aio.insecure_channel") as mock_insecure: + MUT._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + credentials=mock_creds2, + ) + mock_insecure.assert_called_once_with( + target="localhost:9010", interceptors=["single_interceptor"] + ) + + # Credentials with create_auth_interceptor returning list and single item + mock_creds3 = mock.MagicMock(spec=["create_auth_interceptor"]) + mock_creds3.create_auth_interceptor.return_value = ["interceptor3"] + with mock.patch("grpc.aio.insecure_channel") as mock_insecure: + MUT._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + credentials=mock_creds3, + ) + mock_insecure.assert_called_once_with( + target="localhost:9010", interceptors=["interceptor3"] + ) + + mock_creds3.create_auth_interceptor.return_value = "single_interceptor3" + with mock.patch("grpc.aio.insecure_channel") as mock_insecure: + MUT._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + credentials=mock_creds3, + ) + mock_insecure.assert_called_once_with( + target="localhost:9010", interceptors=["single_interceptor3"] + ) + + # TLS and mTLS + with mock.patch("builtins.open", mock.mock_open(read_data=b"cert_data")): + with mock.patch("grpc.ssl_channel_credentials") as mock_ssl_creds: + with mock.patch("grpc.aio.secure_channel") as mock_secure: + # TLS only + MUT._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate=None, + client_key=None, + ) + mock_ssl_creds.assert_called_with(root_certificates=b"cert_data") + mock_secure.assert_called_with( + "omni-host:15000", mock_ssl_creds.return_value, interceptors=[] + ) + + # mTLS + MUT._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate="client.pem", + client_key="key.pem", + ) + mock_ssl_creds.assert_called_with( + root_certificates=b"cert_data", + private_key=b"cert_data", + certificate_chain=b"cert_data", + ) + + # Validation errors + with self.assertRaises(ValueError) as cm: + MUT._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate=None, + client_certificate=None, + client_key=None, + ) + self.assertIn("TLS/mTLS connection requires ca_certificate", str(cm.exception)) + + with mock.patch("builtins.open", mock.mock_open(read_data=b"cert_data")): + with self.assertRaises(ValueError) as cm: + MUT._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate="client.pem", + client_key=None, + ) + self.assertIn( + "Both client_certificate and client_key must be provided for mTLS connection", + str(cm.exception), + ) + + with self.assertRaises(ValueError) as cm: + MUT._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate=None, + client_key="key.pem", + ) + self.assertIn( + "Both client_certificate and client_key must be provided for mTLS connection", + str(cm.exception), + ) diff --git a/packages/google-cloud-spanner/tests/unit/omni/test_credentials.py b/packages/google-cloud-spanner/tests/unit/omni/test_credentials.py new file mode 100644 index 000000000000..9ca506350169 --- /dev/null +++ b/packages/google-cloud-spanner/tests/unit/omni/test_credentials.py @@ -0,0 +1,559 @@ +# -*- 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. + +import base64 +import datetime +import unittest +from collections import namedtuple +from unittest import mock + +import google.auth.exceptions +import grpc +from google.protobuf import timestamp_pb2 + +from google.cloud.spanner_v1.omni.credentials import SpannerOmniCredentials +from google.cloud.spanner_v1.omni.proto import login_pb2 + + +class TestSpannerOmniCredentials(unittest.TestCase): + def test_init_validation(self): + with self.assertRaises(ValueError): + SpannerOmniCredentials("", "password", "localhost:9010") + with self.assertRaises(ValueError): + SpannerOmniCredentials("user", "", "localhost:9010") + with self.assertRaises(ValueError): + SpannerOmniCredentials("user", "password", "") + + def test_target_scheme_parsing(self): + creds_http = SpannerOmniCredentials("user", "pass", "http://localhost:9010") + self.assertEqual(creds_http.target, "localhost:9010") + self.assertTrue(creds_http.use_plain_text) + + creds_https = SpannerOmniCredentials("user", "pass", "https://localhost:9010") + self.assertEqual(creds_https.target, "localhost:9010") + self.assertFalse(creds_https.use_plain_text) + + creds_raw = SpannerOmniCredentials( + "user", "pass", "localhost:9010", use_plain_text=True + ) + self.assertEqual(creds_raw.target, "localhost:9010") + self.assertTrue(creds_raw.use_plain_text) + + def test_refresh_success(self): + creds = SpannerOmniCredentials( + "user", "pass", "localhost:9010", ca_certificate="/dummy/ca.pem" + ) + + mock_token_proto = login_pb2.AccessToken( + username="user", + expiration_time=timestamp_pb2.Timestamp(seconds=1700000000, nanos=0), + signature=b"test_sig", + key_id=42, + ) + + with ( + mock.patch("builtins.open", mock.mock_open(read_data=b"dummy_ca")), + mock.patch("grpc.secure_channel") as mock_sec_channel, + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.return_value = mock_token_proto + + creds.refresh() + + self.assertIsNotNone(creds.token) + # Verify base64 token decodes back to proto + decoded_bytes = base64.b64decode(creds.token) + parsed_token = login_pb2.AccessToken.FromString(decoded_bytes) + self.assertEqual(parsed_token.username, "user") + self.assertEqual(parsed_token.signature, b"test_sig") + self.assertEqual(parsed_token.key_id, 42) + self.assertEqual( + creds.expiry, + datetime.datetime.fromtimestamp( + 1700000000, tz=datetime.timezone.utc + ).replace(tzinfo=None), + ) + mock_sec_channel.return_value.close.assert_called_once() + + def test_refresh_plaintext_channel(self): + creds = SpannerOmniCredentials("user", "pass", "http://localhost:9010") + mock_token_proto = login_pb2.AccessToken(username="user") + + with ( + mock.patch("grpc.insecure_channel") as mock_insec_channel, + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.return_value = mock_token_proto + + creds.refresh() + + mock_insec_channel.assert_called_once_with("localhost:9010") + mock_insec_channel.return_value.close.assert_called_once() + self.assertIsNotNone(creds.token) + + def test_refresh_missing_ca_certificate_raises(self): + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + with self.assertRaises(google.auth.exceptions.RefreshError) as cm: + creds.refresh() + self.assertIn("requires ca_certificate to be set", str(cm.exception)) + + def test_refresh_mtls_success(self): + creds = SpannerOmniCredentials( + "user", + "pass", + "localhost:9010", + ca_certificate="/dummy/ca.pem", + client_certificate="/dummy/cert.pem", + client_key="/dummy/key.pem", + ) + mock_token_proto = login_pb2.AccessToken(username="user") + + with ( + mock.patch("builtins.open", mock.mock_open(read_data=b"dummy_data")), + mock.patch("grpc.ssl_channel_credentials") as mock_ssl_creds, + mock.patch("grpc.secure_channel") as mock_sec_channel, + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.return_value = mock_token_proto + + creds.refresh() + + mock_ssl_creds.assert_called_once_with( + root_certificates=b"dummy_data", + private_key=b"dummy_data", + certificate_chain=b"dummy_data", + ) + mock_sec_channel.assert_called_once_with( + "localhost:9010", mock_ssl_creds.return_value + ) + self.assertIsNotNone(creds.token) + + def test_refresh_zero_expiration_time_falls_back_to_default_ttl(self): + creds = SpannerOmniCredentials( + "user", "pass", "localhost:9010", use_plain_text=True + ) + mock_token_proto = login_pb2.AccessToken(username="user") + mock_token_proto.expiration_time.seconds = 0 + mock_token_proto.expiration_time.nanos = 0 + + with ( + mock.patch("grpc.insecure_channel"), + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.return_value = mock_token_proto + + creds.refresh() + + self.assertIsNotNone(creds.token) + self.assertTrue(creds.valid) + now = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None) + self.assertGreater(creds.expiry, now + datetime.timedelta(minutes=50)) + + def test_refresh_nanos_only_expiration_time(self): + creds = SpannerOmniCredentials( + "user", "pass", "localhost:9010", use_plain_text=True + ) + now_ts = datetime.datetime.now(datetime.timezone.utc).timestamp() + mock_token_proto = login_pb2.AccessToken(username="user") + mock_token_proto.expiration_time.seconds = int(now_ts) + 300 + mock_token_proto.expiration_time.nanos = 500000 + + with ( + mock.patch("grpc.insecure_channel"), + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.return_value = mock_token_proto + + creds.refresh() + + self.assertIsNotNone(creds.token) + self.assertTrue(creds.valid) + now = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None) + self.assertGreater(creds.expiry, now + datetime.timedelta(minutes=4)) + + def test_refresh_mtls_missing_key_raises(self): + creds = SpannerOmniCredentials( + "user", + "pass", + "localhost:9010", + ca_certificate="/dummy/ca.pem", + client_certificate="/dummy/cert.pem", + ) + with ( + mock.patch("builtins.open", mock.mock_open(read_data=b"dummy_data")), + self.assertRaises(google.auth.exceptions.RefreshError) as cm, + ): + creds.refresh() + self.assertIn( + "Both client_certificate and client_key must be provided", + str(cm.exception), + ) + + def test_refresh_failure_wraps_in_refresh_error(self): + creds = SpannerOmniCredentials( + "user", "pass", "localhost:9010", ca_certificate="/dummy/ca.pem" + ) + + with ( + mock.patch("builtins.open", mock.mock_open(read_data=b"dummy_data")), + mock.patch("grpc.secure_channel"), + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.side_effect = ValueError("Handshake failed") + + with self.assertRaises(google.auth.exceptions.RefreshError): + creds.refresh() + + def test_refresh_grpc_error_is_reraised(self): + creds = SpannerOmniCredentials( + "user", "pass", "localhost:9010", ca_certificate="/dummy/ca.pem" + ) + + class CustomRpcError(grpc.RpcError): + pass + + with ( + mock.patch("builtins.open", mock.mock_open(read_data=b"dummy_data")), + mock.patch("grpc.secure_channel"), + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.side_effect = CustomRpcError("Service Unavailable") + + with self.assertRaises(CustomRpcError): + creds.refresh() + + def test_refresh_skips_when_valid_inside_lock(self): + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + creds.token = "existing_valid_token" + creds.expiry = datetime.datetime.now(datetime.timezone.utc).replace( + tzinfo=None + ) + datetime.timedelta(hours=1) + + with mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls: + creds.refresh() + mock_login_client_cls.assert_not_called() + + def test_apply_and_before_request(self): + creds = SpannerOmniCredentials( + "user", "pass", "localhost:9010", ca_certificate="/dummy/ca.pem" + ) + headers = {} + + mock_token_proto = login_pb2.AccessToken(username="user") + with ( + mock.patch("builtins.open", mock.mock_open(read_data=b"dummy_data")), + mock.patch("grpc.secure_channel"), + mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls, + ): + mock_client = mock_login_client_cls.return_value + mock_client.login.return_value = mock_token_proto + + creds.before_request(None, "POST", "http://localhost", headers) + + self.assertIn("authorization", headers) + self.assertTrue(headers["authorization"].startswith("Bearer ")) + + def test_interceptor(self): + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + creds.token = "sample_test_token" + creds.expiry = datetime.datetime.now(datetime.timezone.utc).replace( + tzinfo=None + ) + datetime.timedelta(hours=1) + + interceptor = creds.create_auth_interceptor() + + DummyCallDetails = namedtuple( + "DummyCallDetails", + ["method", "timeout", "metadata", "credentials", "wait_for_ready"], + ) + call_details = DummyCallDetails( + method="/google.spanner.v1.Spanner/ExecuteSql", + timeout=30.0, + metadata=[("custom-header", "custom-val")], + credentials=None, + wait_for_ready=None, + ) + + def dummy_continuation(details, req): + return details + + # Unary-Unary + res = interceptor.intercept_unary_unary(dummy_continuation, call_details, "req") + self.assertIn(("authorization", "Bearer sample_test_token"), res.metadata) + self.assertIn(("custom-header", "custom-val"), res.metadata) + + # Unary-Stream + res_us = interceptor.intercept_unary_stream( + dummy_continuation, call_details, "req" + ) + self.assertIn(("authorization", "Bearer sample_test_token"), res_us.metadata) + + # Stream-Unary + res_su = interceptor.intercept_stream_unary( + dummy_continuation, call_details, ["req"] + ) + self.assertIn(("authorization", "Bearer sample_test_token"), res_su.metadata) + + # Stream-Stream + res_ss = interceptor.intercept_stream_stream( + dummy_continuation, call_details, ["req"] + ) + self.assertIn(("authorization", "Bearer sample_test_token"), res_ss.metadata) + + def test_async_interceptor_attaches_bearer_token(self): + import asyncio + + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + creds.token = "async_test_token" + creds.expiry = datetime.datetime.now(datetime.timezone.utc).replace( + tzinfo=None + ) + datetime.timedelta(hours=1) + + interceptors = creds.create_async_auth_interceptors() + self.assertEqual(len(interceptors), 4) + + DummyCallDetails = namedtuple( + "DummyCallDetails", + ["method", "timeout", "metadata", "credentials", "wait_for_ready"], + ) + call_details = DummyCallDetails( + method="/google.spanner.v1.Spanner/StreamingRead", + timeout=30.0, + metadata=[("custom-header", "custom-val")], + credentials=None, + wait_for_ready=None, + ) + + async def dummy_async_continuation(details, req): + return details + + async def run_async_tests(): + # Unary-Unary + res_uu = await interceptors[0].intercept_unary_unary( + dummy_async_continuation, call_details, "req" + ) + self.assertIn(("authorization", "Bearer async_test_token"), res_uu.metadata) + self.assertIn(("custom-header", "custom-val"), res_uu.metadata) + + # Unary-Stream + res_us = await interceptors[1].intercept_unary_stream( + dummy_async_continuation, call_details, "req" + ) + self.assertIn(("authorization", "Bearer async_test_token"), res_us.metadata) + + # Stream-Unary + res_su = await interceptors[2].intercept_stream_unary( + dummy_async_continuation, call_details, ["req"] + ) + self.assertIn(("authorization", "Bearer async_test_token"), res_su.metadata) + + # Stream-Stream + res_ss = await interceptors[3].intercept_stream_stream( + dummy_async_continuation, call_details, ["req"] + ) + self.assertIn(("authorization", "Bearer async_test_token"), res_ss.metadata) + + asyncio.run(run_async_tests()) + + def test_async_interceptor_triggers_refresh_when_invalid(self): + import asyncio + + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + interceptors = creds.create_async_auth_interceptors() + + DummyCallDetails = namedtuple( + "DummyCallDetails", + ["method", "timeout", "metadata", "credentials", "wait_for_ready"], + ) + call_details = DummyCallDetails( + method="/google.spanner.v1.Spanner/ExecuteSql", + timeout=30.0, + metadata=[], + credentials=None, + wait_for_ready=None, + ) + + async def dummy_async_continuation(details, req): + return details + + async def run_test(): + with mock.patch.object(creds, "refresh") as mock_refresh: + + def do_refresh(): + creds.token = "refreshed_async_token" + creds.expiry = datetime.datetime.now(datetime.timezone.utc).replace( + tzinfo=None + ) + datetime.timedelta(hours=1) + + mock_refresh.side_effect = do_refresh + + res = await interceptors[0].intercept_unary_unary( + dummy_async_continuation, call_details, "req" + ) + mock_refresh.assert_called_once() + self.assertIn( + ("authorization", "Bearer refreshed_async_token"), res.metadata + ) + + asyncio.run(run_test()) + + def test_init_channel(self): + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + mock_ssl = mock.Mock(spec=grpc.ChannelCredentials) + creds.init_channel( + use_plain_text=True, + ca_certificate="ca.pem", + client_certificate="client.pem", + client_key="key.pem", + ssl_credentials=mock_ssl, + ) + self.assertTrue(creds.use_plain_text) + self.assertEqual(creds.ca_certificate, "ca.pem") + self.assertEqual(creds.client_certificate, "client.pem") + self.assertEqual(creds.client_key, "key.pem") + self.assertIs(creds.ssl_credentials, mock_ssl) + + # Also test with use_plain_text=False + creds.init_channel(use_plain_text=False) + self.assertFalse(creds.use_plain_text) + + def test_create_auth_interceptor_async_dispatch(self): + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + interceptors = creds.create_auth_interceptor(is_async=True) + self.assertIsInstance(interceptors, list) + self.assertEqual(len(interceptors), 4) + + interceptors_alias = creds.create_async_auth_interceptor() + self.assertIsInstance(interceptors_alias, list) + self.assertEqual(len(interceptors_alias), 4) + + def test_sync_interceptor_triggers_refresh_when_invalid(self): + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + interceptor = creds.create_auth_interceptor() + + DummyCallDetails = namedtuple( + "DummyCallDetails", + ["method", "timeout", "metadata", "credentials", "wait_for_ready"], + ) + call_details = DummyCallDetails( + method="/google.spanner.v1.Spanner/ExecuteSql", + timeout=30.0, + metadata=[], + credentials=None, + wait_for_ready=None, + ) + + def dummy_continuation(details, req): + return details + + with mock.patch.object(creds, "refresh") as mock_refresh: + + def do_refresh(): + creds.token = "refreshed_sync_token" + creds.expiry = datetime.datetime.now(datetime.timezone.utc).replace( + tzinfo=None + ) + datetime.timedelta(hours=1) + + mock_refresh.side_effect = do_refresh + + res = interceptor.intercept_unary_unary( + dummy_continuation, call_details, "req" + ) + mock_refresh.assert_called_once() + self.assertIn( + ("authorization", "Bearer refreshed_sync_token"), res.metadata + ) + + def test_perform_refresh_token_with_ssl_credentials(self): + mock_ssl = mock.Mock(spec=grpc.ChannelCredentials) + creds = SpannerOmniCredentials( + "user", + "pass", + "localhost:9010", + ssl_credentials=mock_ssl, + ) + with mock.patch("grpc.secure_channel") as mock_secure_channel: + mock_channel = mock.MagicMock() + mock_secure_channel.return_value = mock_channel + with mock.patch( + "google.cloud.spanner_v1.omni.credentials.LoginClient" + ) as mock_login_client_cls: + mock_client = mock.MagicMock() + mock_login_client_cls.return_value = mock_client + proto_token = login_pb2.AccessToken( + username="user", signature=b"secure_token" + ) + mock_client.login.return_value = proto_token + + creds.refresh() + + mock_secure_channel.assert_called_once_with("localhost:9010", mock_ssl) + self.assertTrue(creds.valid) + self.assertIsNotNone(creds.token) + mock_channel.close.assert_called_once() + + def test_before_request(self): + creds = SpannerOmniCredentials("user", "pass", "localhost:9010") + headers = {} + + with mock.patch.object(creds, "refresh") as mock_refresh: + + def do_refresh(request=None): + creds.token = "refreshed_token" + creds.expiry = datetime.datetime.now(datetime.timezone.utc).replace( + tzinfo=None + ) + datetime.timedelta(hours=1) + + mock_refresh.side_effect = do_refresh + + creds.before_request(None, "GET", "http://example.com", headers) + mock_refresh.assert_called_once() + self.assertEqual(headers["authorization"], "Bearer refreshed_token") + + # Call again when valid - refresh should not be called again + mock_refresh.reset_mock() + headers_2 = {} + creds.before_request(None, "GET", "http://example.com", headers_2) + mock_refresh.assert_not_called() + self.assertEqual(headers_2["authorization"], "Bearer refreshed_token") + + +if __name__ == "__main__": + unittest.main() diff --git a/packages/google-cloud-spanner/tests/unit/omni/test_login_client.py b/packages/google-cloud-spanner/tests/unit/omni/test_login_client.py new file mode 100644 index 000000000000..3d080cd67859 --- /dev/null +++ b/packages/google-cloud-spanner/tests/unit/omni/test_login_client.py @@ -0,0 +1,512 @@ +# -*- 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. + +import unittest +from unittest import mock + +import grpc +from google.protobuf import timestamp_pb2 + +from google.cloud.spanner_v1.omni import opaque +from google.cloud.spanner_v1.omni.login_client import LoginClient +from google.cloud.spanner_v1.omni.proto import authentication_pb2, login_pb2 + + +class TestLoginClient(unittest.TestCase): + def setUp(self): + self.mock_channel = mock.MagicMock(spec=grpc.Channel) + + def test_login_empty_inputs(self): + client = LoginClient(self.mock_channel) + with self.assertRaises(ValueError): + client.login("", "password") + with self.assertRaises(ValueError): + client.login("user", "") + + def test_login_successful_flow(self): + username = "admin" + password = b"secret123" + + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + + # Precompute server artifacts + oprf_seed = opaque.nonce() + oprf_key_seed = opaque.expand( + oprf_seed, (username + "OprfKey").encode("utf-8"), 32 + ) + _, oprf_priv = opaque.derive_key_pair(oprf_key_seed, b"OPAQUE-DeriveKeyPair") + + h_pt = opaque.hash_to_curve_p256(password, opaque.LOGIN_DOMAIN_SEPARATION_TAG) + prf_pt = opaque.point_mul(h_pt, int.from_bytes(oprf_priv, "big")) + prf = opaque.marshal_compressed(prf_pt) + + stretched_oprf = opaque.stretch(prf, params) + randomized_password = opaque.extract(opaque.concat(prf, stretched_oprf)) + + server_key_seed = opaque.nonce() + server_pub, server_priv = opaque.derive_key_pair( + server_key_seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + + envelope_nonce = opaque.nonce() + auth_key = opaque.expand( + randomized_password, envelope_nonce + opaque.AUTH_KEY_INFO, 32 + ) + auth_tag = opaque.mac( + auth_key, envelope_nonce + server_pub + username.encode("utf-8") + ) + serialized_envelope = opaque.concat(server_pub, envelope_nonce, auth_tag) + + masking_key = opaque.expand(randomized_password, opaque.MASKING_KEY_INFO, 32) + masking_nonce = opaque.nonce() + credential_pad = opaque.expand( + masking_key, + opaque.concat(masking_nonce, b"CredentialResponsePad"), + len(serialized_envelope), + ) + masked_response = opaque.xor_bytes(serialized_envelope, credential_pad) + + expected_access_token = login_pb2.AccessToken( + username=username, + creation_time=timestamp_pb2.Timestamp(seconds=1000, nanos=0), + expiration_time=timestamp_pb2.Timestamp(seconds=4600, nanos=0), + signature=b"valid_signature", + key_id=1, + access_token_type=login_pb2.AccessToken.ACCESS_TOKEN_TYPE_API, + ) + + def mock_login_rpc(request_iterator, timeout=None): + # Step 1: receive handshake request + req1 = next(request_iterator) + self.assertEqual(req1.username, username) + self.assertTrue(req1.HasField("handshake_request")) + + # Return Step 1 response + yield login_pb2.LoginResponse( + handshake_response=authentication_pb2.PasswordAuthenticationHandshakeResponse( + password_authentication_protocol=authentication_pb2.PasswordAuthenticationProtocol.PASSWORD_AUTHENTICATION_PROTOCOL_OPAQUE, + hash_parameters=params, + ) + ) + + # Step 2: receive initial opaque request + req2 = next(request_iterator) + init_req = req2.opaque_request.initial_request + blinded_msg = init_req.blinded_message + client_nonce = init_req.client_nonce + client_pub_keyshare = init_req.client_public_keyshare + + blinded_pt = opaque.unmarshal_compressed(blinded_msg) + eval_pt = opaque.point_mul(blinded_pt, int.from_bytes(oprf_priv, "big")) + evaluated_msg = opaque.marshal_compressed(eval_pt) + + server_ephemeral_pub, server_ephemeral_priv = opaque.derive_key_pair( + opaque.nonce(), opaque.DIFFIE_HELLMAN_KEY_INFO + ) + server_login_nonce = opaque.nonce() + + seed = opaque.expand( + randomized_password, envelope_nonce + opaque.PRIVATE_KEY_INFO, 32 + ) + client_pub, _ = opaque.derive_key_pair(seed, opaque.DIFFIE_HELLMAN_KEY_INFO) + + s_dh1 = opaque.diffie_hellman(server_ephemeral_priv, client_pub_keyshare) + s_dh2 = opaque.diffie_hellman(server_priv, client_pub_keyshare) + s_dh3 = opaque.diffie_hellman(server_ephemeral_priv, client_pub) + s_ikm = opaque.concat(s_dh1, s_dh2, s_dh3) + + preamble = opaque.concat( + b"OPAQUEv1-", + username.encode("utf-8"), + client_nonce, + client_pub_keyshare, + server_pub, + evaluated_msg, + server_login_nonce, + server_ephemeral_pub, + ) + + s_km2, _, _ = opaque.derive_shared_keys(s_ikm, preamble) + server_mac = opaque.mac(s_km2, opaque.sha256_hash(preamble)) + + # Return Step 2 response + yield login_pb2.LoginResponse( + opaque_response=login_pb2.OpaqueLoginResponse( + initial_response=login_pb2.InitialOpaqueLoginResponse( + server_nonce=server_login_nonce, + server_public_keyshare=server_ephemeral_pub, + server_mac=server_mac, + evaluated_message=evaluated_msg, + masking_nonce=masking_nonce, + masked_response=masked_response, + ) + ) + ) + + # Step 3: receive final opaque request + req3 = next(request_iterator) + self.assertTrue(req3.opaque_request.HasField("final_request")) + + # Return Step 3 response + yield login_pb2.LoginResponse(access_token=expected_access_token) + + with mock.patch( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc.LoginServiceStub" + ) as mock_stub_cls: + mock_stub = mock_stub_cls.return_value + mock_stub.Login.side_effect = mock_login_rpc + + client = LoginClient(self.mock_channel) + token = client.login(username, password) + + self.assertEqual(token.username, username) + self.assertEqual(token.signature, b"valid_signature") + self.assertEqual(token.key_id, 1) + self.assertEqual(token.expiration_time.seconds, 4600) + + def test_login_unsupported_protocol(self): + def mock_login_rpc(request_iterator, timeout=None): + next(request_iterator) + yield login_pb2.LoginResponse( + handshake_response=authentication_pb2.PasswordAuthenticationHandshakeResponse( + password_authentication_protocol=authentication_pb2.PasswordAuthenticationProtocol.PASSWORD_AUTHENTICATION_PROTOCOL_UNSPECIFIED, + ) + ) + + with mock.patch( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc.LoginServiceStub" + ) as mock_stub_cls: + mock_stub = mock_stub_cls.return_value + mock_stub.Login.side_effect = mock_login_rpc + + client = LoginClient(self.mock_channel) + with self.assertRaises(ValueError) as cm: + client.login("user", "pass") + self.assertIn( + "Unsupported password authentication protocol", str(cm.exception) + ) + + def test_login_missing_handshake_response(self): + def mock_login_rpc(request_iterator, timeout=None): + next(request_iterator) + yield login_pb2.LoginResponse() + + with mock.patch( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc.LoginServiceStub" + ) as mock_stub_cls: + mock_stub = mock_stub_cls.return_value + mock_stub.Login.side_effect = mock_login_rpc + + client = LoginClient(self.mock_channel) + with self.assertRaises(ValueError) as cm: + client.login("user", "pass") + self.assertIn("Failed to receive handshake response", str(cm.exception)) + + def test_login_missing_hash_parameters(self): + def mock_login_rpc(request_iterator, timeout=None): + next(request_iterator) + yield login_pb2.LoginResponse( + handshake_response=authentication_pb2.PasswordAuthenticationHandshakeResponse( + password_authentication_protocol=authentication_pb2.PasswordAuthenticationProtocol.PASSWORD_AUTHENTICATION_PROTOCOL_OPAQUE, + ) + ) + + with mock.patch( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc.LoginServiceStub" + ) as mock_stub_cls: + mock_stub = mock_stub_cls.return_value + mock_stub.Login.side_effect = mock_login_rpc + + client = LoginClient(self.mock_channel) + with self.assertRaises(ValueError) as cm: + client.login("user", "pass") + self.assertIn( + "Handshake response missing hash_parameters", str(cm.exception) + ) + + def test_login_missing_access_token_in_final_response(self): + username = "test_user" + password = "test_password" + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + + oprf_seed = opaque.nonce() + oprf_key_seed = opaque.expand( + oprf_seed, (username + "OprfKey").encode("utf-8"), 32 + ) + _, oprf_priv = opaque.derive_key_pair(oprf_key_seed, b"OPAQUE-DeriveKeyPair") + h_pt = opaque.hash_to_curve_p256( + password.encode("utf-8"), opaque.LOGIN_DOMAIN_SEPARATION_TAG + ) + prf_pt = opaque.point_mul(h_pt, int.from_bytes(oprf_priv, "big")) + prf = opaque.marshal_compressed(prf_pt) + stretched_oprf = opaque.stretch(prf, params) + randomized_password = opaque.extract(opaque.concat(prf, stretched_oprf)) + + server_key_seed = opaque.nonce() + server_pub, server_priv = opaque.derive_key_pair( + server_key_seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + envelope_nonce = opaque.nonce() + auth_key = opaque.expand( + randomized_password, envelope_nonce + opaque.AUTH_KEY_INFO, 32 + ) + auth_tag = opaque.mac( + auth_key, envelope_nonce + server_pub + username.encode("utf-8") + ) + serialized_envelope = opaque.concat(server_pub, envelope_nonce, auth_tag) + masking_key = opaque.expand(randomized_password, opaque.MASKING_KEY_INFO, 32) + masking_nonce = opaque.nonce() + credential_pad = opaque.expand( + masking_key, + opaque.concat(masking_nonce, b"CredentialResponsePad"), + len(serialized_envelope), + ) + masked_response = opaque.xor_bytes(serialized_envelope, credential_pad) + + def mock_login_rpc(request_iterator, timeout=None): + next(request_iterator) + yield login_pb2.LoginResponse( + handshake_response=authentication_pb2.PasswordAuthenticationHandshakeResponse( + password_authentication_protocol=authentication_pb2.PasswordAuthenticationProtocol.PASSWORD_AUTHENTICATION_PROTOCOL_OPAQUE, + hash_parameters=params, + ) + ) + + req2 = next(request_iterator) + init_req = req2.opaque_request.initial_request + blinded_msg = init_req.blinded_message + client_nonce = init_req.client_nonce + client_pub_keyshare = init_req.client_public_keyshare + + blinded_pt = opaque.unmarshal_compressed(blinded_msg) + eval_pt = opaque.point_mul(blinded_pt, int.from_bytes(oprf_priv, "big")) + evaluated_msg = opaque.marshal_compressed(eval_pt) + + server_ephemeral_pub, server_ephemeral_priv = opaque.derive_key_pair( + opaque.nonce(), opaque.DIFFIE_HELLMAN_KEY_INFO + ) + server_login_nonce = opaque.nonce() + seed = opaque.expand( + randomized_password, envelope_nonce + opaque.PRIVATE_KEY_INFO, 32 + ) + client_pub, _ = opaque.derive_key_pair(seed, opaque.DIFFIE_HELLMAN_KEY_INFO) + + s_dh1 = opaque.diffie_hellman(server_ephemeral_priv, client_pub_keyshare) + s_dh2 = opaque.diffie_hellman(server_priv, client_pub_keyshare) + s_dh3 = opaque.diffie_hellman(server_ephemeral_priv, client_pub) + s_ikm = opaque.concat(s_dh1, s_dh2, s_dh3) + + preamble = opaque.concat( + b"OPAQUEv1-", + username.encode("utf-8"), + client_nonce, + client_pub_keyshare, + server_pub, + evaluated_msg, + server_login_nonce, + server_ephemeral_pub, + ) + s_km2, _, _ = opaque.derive_shared_keys(s_ikm, preamble) + server_mac = opaque.mac(s_km2, opaque.sha256_hash(preamble)) + + yield login_pb2.LoginResponse( + opaque_response=login_pb2.OpaqueLoginResponse( + initial_response=login_pb2.InitialOpaqueLoginResponse( + server_nonce=server_login_nonce, + server_public_keyshare=server_ephemeral_pub, + server_mac=server_mac, + evaluated_message=evaluated_msg, + masking_nonce=masking_nonce, + masked_response=masked_response, + ) + ) + ) + + next(request_iterator) + # Response without access_token + yield login_pb2.LoginResponse() + + with mock.patch( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc.LoginServiceStub" + ) as mock_stub_cls: + mock_stub = mock_stub_cls.return_value + mock_stub.Login.side_effect = mock_login_rpc + + client = LoginClient(self.mock_channel) + with self.assertRaises(ValueError) as cm: + client.login(username, password) + self.assertIn( + "Server failed to return an access token in final response", + str(cm.exception), + ) + + def test_request_iterator_iter_and_stop(self): + from google.cloud.spanner_v1.omni.login_client import _RequestIterator + + req_iterator = _RequestIterator() + self.assertIs(iter(req_iterator), req_iterator) + + req = login_pb2.LoginRequest(username="u") + req_iterator.send(req) + self.assertEqual(next(req_iterator), req) + + req_iterator.close() + with self.assertRaises(StopIteration): + next(req_iterator) + + def test_request_iterator_close_idempotent(self): + from google.cloud.spanner_v1.omni.login_client import _RequestIterator + + req_iterator = _RequestIterator() + req_iterator.close() + req_iterator.close() + self.assertEqual(req_iterator._queue.qsize(), 1) + self.assertIsNone(req_iterator._queue.get()) + self.assertTrue(req_iterator._queue.empty()) + + def test_login_exception_cancels_call(self): + mock_call = mock.MagicMock() + mock_call.__next__.side_effect = RuntimeError("network error") + + with mock.patch( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc.LoginServiceStub" + ) as mock_stub_cls: + mock_stub = mock_stub_cls.return_value + mock_stub.Login.return_value = mock_call + + client = LoginClient(self.mock_channel) + with self.assertRaises(RuntimeError): + client.login("user", "pass") + + mock_call.cancel.assert_called_once() + + def test_login_premature_stream_closure(self): + mock_call = mock.MagicMock() + mock_call.__next__.side_effect = StopIteration + + with mock.patch( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc.LoginServiceStub" + ) as mock_stub_cls: + mock_stub = mock_stub_cls.return_value + mock_stub.Login.return_value = mock_call + + client = LoginClient(self.mock_channel) + with self.assertRaises(ValueError) as cm: + client.login("user", "pass") + self.assertIn( + "Server closed stream prematurely during handshake", str(cm.exception) + ) + + def test_grpc_servicers_and_helpers(self): + import grpc + + from google.cloud.spanner_v1.omni.proto import login_pb2_grpc + + servicer = login_pb2_grpc.LoginServiceServicer() + mock_context = mock.MagicMock() + with self.assertRaises(NotImplementedError): + servicer.Login(iter([]), mock_context) + mock_context.set_code.assert_called_once_with(grpc.StatusCode.UNIMPLEMENTED) + mock_context.set_details.assert_called_once_with("Method not implemented!") + + mock_server = mock.MagicMock() + login_pb2_grpc.add_LoginServiceServicer_to_server(servicer, mock_server) + mock_server.add_generic_rpc_handlers.assert_called_once() + mock_server.add_registered_method_handlers.assert_called_once() + + with mock.patch("grpc.experimental.stream_stream") as mock_ss: + login_pb2_grpc.LoginService.Login(iter([]), "target_host") + mock_ss.assert_called_once() + + def test_grpc_version_mismatch_paths(self): + import importlib + + try: + importlib.import_module("grpc._utilities") + except ImportError: + pass + + # Test login_pb2_grpc version mismatch check + with mock.patch( + "grpc._utilities.first_version_is_lower", create=True, return_value=True + ): + with self.assertRaises(RuntimeError) as cm: + importlib.reload( + importlib.import_module( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc" + ) + ) + self.assertIn("The grpc package installed is at version", str(cm.exception)) + + with mock.patch( + "grpc._utilities.first_version_is_lower", create=True, return_value=True + ): + with self.assertRaises(RuntimeError) as cm: + importlib.reload( + importlib.import_module( + "google.cloud.spanner_v1.omni.proto.authentication_pb2_grpc" + ) + ) + self.assertIn("The grpc package installed is at version", str(cm.exception)) + + # Test ImportError fallback when first_version_is_lower cannot be imported + import sys + + orig_utilities = sys.modules.get("grpc._utilities") + try: + sys.modules["grpc._utilities"] = None + importlib.reload( + importlib.import_module( + "google.cloud.spanner_v1.omni.proto.login_pb2_grpc" + ) + ) + importlib.reload( + importlib.import_module( + "google.cloud.spanner_v1.omni.proto.authentication_pb2_grpc" + ) + ) + finally: + if orig_utilities is not None: + sys.modules["grpc._utilities"] = orig_utilities + else: + sys.modules.pop("grpc._utilities", None) + + # Restore modules to standard state + importlib.reload( + importlib.import_module("google.cloud.spanner_v1.omni.proto.login_pb2_grpc") + ) + importlib.reload( + importlib.import_module( + "google.cloud.spanner_v1.omni.proto.authentication_pb2_grpc" + ) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/packages/google-cloud-spanner/tests/unit/omni/test_opaque.py b/packages/google-cloud-spanner/tests/unit/omni/test_opaque.py new file mode 100644 index 000000000000..894a6a3b22b0 --- /dev/null +++ b/packages/google-cloud-spanner/tests/unit/omni/test_opaque.py @@ -0,0 +1,1020 @@ +# -*- 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. + +import unittest + +from google.cloud.spanner_v1.omni import opaque +from google.cloud.spanner_v1.omni.proto import authentication_pb2, login_pb2 + + +class TestOpaqueCrypto(unittest.TestCase): + def test_random_oracle_sha256(self): + max_val = 1 << 63 + test_inputs = [b"key", b"key2", bytes([97, 97, 98, 99, 100, 101])] + for inp in test_inputs: + expected = opaque.random_oracle_sha256(inp, max_val) + self.assertEqual(len(expected), 32) + for _ in range(10): + out = opaque.random_oracle_sha256(inp, max_val) + self.assertEqual(out, expected) + + def test_random_oracle_sha256_large_domain_raises(self): + max_val = 1 << 65500 + with self.assertRaises(ValueError) as cm: + opaque.random_oracle_sha256(b"key", max_val) + self.assertIn( + "Domain bit length must not be greater than 65280", str(cm.exception) + ) + + def test_mac(self): + tests = [ + (b"key", b"data"), + (b"key", b"data2"), + ( + bytes([97, 97, 98, 99, 100, 101]), + bytes([102, 103, 104, 105, 106, 107]), + ), + ] + for k, d in tests: + m1 = opaque.mac(k, d) + m2 = opaque.mac(k, d) + self.assertEqual(m1, m2) + self.assertEqual(len(m1), 32) + + def test_xor_bytes(self): + tests = [ + (b"abc", b"def", False), + ( + bytes([97, 97, 98, 99, 100, 101]), + bytes([102, 103, 104, 105, 106, 107]), + False, + ), + ( + bytes([97, 97, 98, 99, 100, 101]), + bytes([0, 0, 0, 0, 0, 0]), + False, + ), + (b"abc", b"defghi", True), + (b"abcdefghi", b"jklmnop", True), + (b"", b"", False), + ] + for a, b, want_err in tests: + if want_err: + with self.assertRaises(ValueError): + opaque.xor_bytes(a, b) + else: + xored = opaque.xor_bytes(a, b) + self.assertEqual(len(xored), len(a)) + orig = opaque.xor_bytes(xored, b) + self.assertEqual(orig, a) + + def test_stretch(self): + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=5, + memory_usage=7 * 1024, + parallelism=1, + hash_size=32, + ) + ) + long_input = bytes(range(256)) * 4 + + tests = [ + ( + b"", + bytes( + [ + 58, + 42, + 135, + 162, + 54, + 231, + 153, + 103, + 111, + 241, + 220, + 39, + 245, + 158, + 231, + 5, + 157, + 108, + 133, + 178, + 37, + 97, + 185, + 220, + 104, + 13, + 66, + 147, + 221, + 19, + 198, + 9, + ] + ), + ), + ( + b"input", + bytes( + [ + 177, + 173, + 204, + 142, + 245, + 214, + 91, + 164, + 139, + 85, + 150, + 101, + 204, + 187, + 48, + 176, + 251, + 7, + 154, + 247, + 251, + 35, + 241, + 135, + 99, + 117, + 14, + 121, + 182, + 124, + 87, + 46, + ] + ), + ), + ( + bytes([97, 97, 98, 99, 100, 101]), + bytes( + [ + 164, + 94, + 8, + 109, + 17, + 19, + 42, + 55, + 86, + 44, + 54, + 89, + 255, + 148, + 130, + 248, + 133, + 4, + 40, + 24, + 246, + 27, + 81, + 56, + 231, + 137, + 238, + 30, + 67, + 159, + 3, + 157, + ] + ), + ), + ( + long_input, + bytes( + [ + 132, + 52, + 182, + 135, + 97, + 18, + 8, + 254, + 10, + 1, + 94, + 98, + 78, + 193, + 246, + 160, + 12, + 209, + 142, + 253, + 247, + 115, + 4, + 149, + 141, + 2, + 105, + 159, + 139, + 94, + 161, + 116, + ] + ), + ), + ] + for inp, expected in tests: + stretched = opaque.stretch(inp, params) + self.assertEqual(len(stretched), 32) + self.assertEqual(stretched, expected) + + def test_extract(self): + long_input = bytes(range(256)) * 4 + tests = [ + ( + b"", + bytes( + [ + 99, + 252, + 241, + 111, + 84, + 209, + 178, + 181, + 88, + 96, + 91, + 194, + 149, + 79, + 240, + 143, + 252, + 68, + 135, + 177, + 69, + 144, + 33, + 115, + 195, + 224, + 100, + 31, + 46, + 160, + 150, + 41, + ] + ), + ), + ( + b"input", + bytes( + [ + 94, + 113, + 123, + 114, + 170, + 250, + 213, + 241, + 247, + 203, + 160, + 141, + 111, + 233, + 68, + 240, + 123, + 33, + 207, + 139, + 115, + 44, + 249, + 217, + 77, + 34, + 6, + 254, + 77, + 75, + 20, + 99, + ] + ), + ), + ( + bytes([97, 97, 98, 99, 100, 101]), + bytes( + [ + 48, + 112, + 244, + 9, + 53, + 2, + 10, + 147, + 218, + 132, + 43, + 198, + 200, + 101, + 20, + 3, + 71, + 158, + 227, + 3, + 161, + 15, + 215, + 112, + 251, + 195, + 187, + 96, + 11, + 203, + 226, + 210, + ] + ), + ), + ( + long_input, + bytes( + [ + 246, + 148, + 220, + 16, + 96, + 62, + 53, + 189, + 96, + 83, + 146, + 84, + 233, + 183, + 89, + 12, + 235, + 31, + 24, + 113, + 148, + 25, + 213, + 33, + 167, + 78, + 147, + 162, + 223, + 115, + 38, + 117, + ] + ), + ), + ] + for inp, expected in tests: + extracted = opaque.extract(inp) + self.assertEqual(len(extracted), 32) + self.assertEqual(extracted, expected) + + def test_derive_key_pair(self): + tests = [ + (b"seed", b"info", b"seed", b"info", False), + (b"seed2", b"info", b"seed2", b"info", False), + (b"seed", b"info2", b"seed", b"info2", False), + (b"seed", b"info2", b"different", b"info2", True), + (b"seed", b"info2", b"seed", b"info1", True), + ] + for s1, i1, s2, i2, want_diff in tests: + pub1, priv1 = opaque.derive_key_pair(s1, i1) + pub2, priv2 = opaque.derive_key_pair(s2, i2) + self.assertEqual(len(pub1), 33) + self.assertEqual(len(priv1), 32) + if want_diff: + self.assertNotEqual(priv1, priv2) + self.assertNotEqual(pub1, pub2) + else: + self.assertEqual(priv1, priv2) + self.assertEqual(pub1, pub2) + + def test_diffie_hellman(self): + tests = [ + (b"", b""), + (b"server-seed", b"client-seed"), + (b"server-seed2", b"client-seed2"), + (b"no-need-to-be-the-same-length", b"im-a-shorter-seed"), + ] + for server_seed, client_seed in tests: + server_pub, server_priv = opaque.derive_key_pair( + server_seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + client_pub, client_priv = opaque.derive_key_pair( + client_seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + + server_shared = opaque.diffie_hellman(server_priv, client_pub) + client_shared = opaque.diffie_hellman(client_priv, server_pub) + self.assertEqual(server_shared, client_shared) + + def test_oprf_evaluate(self): + username = "username" + password = b"password1234" + oprf_seed = opaque.nonce() + seed = opaque.expand(oprf_seed, (username + "OprfKey").encode("utf-8"), 32) + _, server_priv = opaque.derive_key_pair(seed, b"OPAQUE-DeriveKeyPair") + + blinded_element, blind_scalar = opaque.blind(password) + + # Server blind evaluation + blinded_pt = opaque.unmarshal_compressed(blinded_element) + evaluated_pt = opaque.point_mul(blinded_pt, int.from_bytes(server_priv, "big")) + evaluated_element = opaque.marshal_compressed(evaluated_pt) + + # Client finalization + oprf = opaque.finalize(blind_scalar, evaluated_element) + + # Direct evaluation + h_pt = opaque.hash_to_curve_p256(password, opaque.LOGIN_DOMAIN_SEPARATION_TAG) + prf_pt = opaque.point_mul(h_pt, int.from_bytes(server_priv, "big")) + prf = opaque.marshal_compressed(prf_pt) + + self.assertEqual(oprf, prf) + + def test_authenticator_validation(self): + valid_params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + + with self.assertRaises(ValueError): + opaque.UserAuthenticator("", b"pass", valid_params) + + with self.assertRaises(ValueError): + opaque.UserAuthenticator("user", b"", valid_params) + + with self.assertRaises(ValueError): + opaque.UserAuthenticator("user", b"pass", None) + + with self.assertRaises(ValueError): + opaque.UserAuthenticator( + "user", + b"pass", + authentication_pb2.HashParameters(), + ).initial_request() + + def test_authenticator_state_errors(self): + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + auth = opaque.UserAuthenticator("user", "password", params) + + # Final before initial + resp = login_pb2.LoginResponse( + opaque_response=login_pb2.OpaqueLoginResponse( + initial_response=login_pb2.InitialOpaqueLoginResponse() + ) + ) + with self.assertRaises(ValueError): + auth.final_request(resp) + + # First initial works + req1 = auth.initial_request() + self.assertEqual(req1.username, "user") + self.assertTrue(req1.opaque_request.HasField("initial_request")) + + # Second initial fails + with self.assertRaises(ValueError): + auth.initial_request() + + def test_full_opaque_handshake_simulation(self): + username = "alice" + password = b"secret_password_123" + + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + + # 1. Server setup user registration: + # oprf_key + oprf_seed = opaque.nonce() + oprf_key_seed = opaque.expand( + oprf_seed, (username + "OprfKey").encode("utf-8"), 32 + ) + _, oprf_priv = opaque.derive_key_pair(oprf_key_seed, b"OPAQUE-DeriveKeyPair") + + # Direct prf for enrollment + h_pt = opaque.hash_to_curve_p256(password, opaque.LOGIN_DOMAIN_SEPARATION_TAG) + prf_pt = opaque.point_mul(h_pt, int.from_bytes(oprf_priv, "big")) + prf = opaque.marshal_compressed(prf_pt) + + stretched_oprf = opaque.stretch(prf, params) + randomized_password = opaque.extract(opaque.concat(prf, stretched_oprf)) + + # Server keypair + server_key_seed = opaque.nonce() + server_pub, server_priv = opaque.derive_key_pair( + server_key_seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + + # Client keypair & envelope + envelope_nonce = opaque.nonce() + auth_key = opaque.expand( + randomized_password, envelope_nonce + opaque.AUTH_KEY_INFO, 32 + ) + auth_tag = opaque.mac( + auth_key, envelope_nonce + server_pub + username.encode("utf-8") + ) + serialized_envelope = opaque.concat(server_pub, envelope_nonce, auth_tag) + + masking_key = opaque.expand(randomized_password, opaque.MASKING_KEY_INFO, 32) + masking_nonce = opaque.nonce() + credential_pad = opaque.expand( + masking_key, + opaque.concat(masking_nonce, b"CredentialResponsePad"), + len(serialized_envelope), + ) + masked_response = opaque.xor_bytes(serialized_envelope, credential_pad) + + # 2. Client starts login + client_auth = opaque.UserAuthenticator(username, password, params) + client_req1 = client_auth.initial_request() + + blinded_msg = client_req1.opaque_request.initial_request.blinded_message + client_nonce = client_req1.opaque_request.initial_request.client_nonce + client_pub_keyshare = ( + client_req1.opaque_request.initial_request.client_public_keyshare + ) + + # 3. Server processes initial request: + blinded_pt = opaque.unmarshal_compressed(blinded_msg) + eval_pt = opaque.point_mul(blinded_pt, int.from_bytes(oprf_priv, "big")) + evaluated_msg = opaque.marshal_compressed(eval_pt) + + server_ephemeral_seed = opaque.nonce() + server_ephemeral_pub, server_ephemeral_priv = opaque.derive_key_pair( + server_ephemeral_seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + server_login_nonce = opaque.nonce() + + # Server computes shared keys + # Client recovered pub key from envelope + seed = opaque.expand( + randomized_password, envelope_nonce + opaque.PRIVATE_KEY_INFO, 32 + ) + client_pub, client_priv = opaque.derive_key_pair( + seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + + s_dh1 = opaque.diffie_hellman(server_ephemeral_priv, client_pub_keyshare) + s_dh2 = opaque.diffie_hellman(server_priv, client_pub_keyshare) + s_dh3 = opaque.diffie_hellman(server_ephemeral_priv, client_pub) + s_ikm = opaque.concat(s_dh1, s_dh2, s_dh3) + + preamble = opaque.concat( + b"OPAQUEv1-", + username.encode("utf-8"), + client_nonce, + client_pub_keyshare, + server_pub, + evaluated_msg, + server_login_nonce, + server_ephemeral_pub, + ) + + s_km2, s_km3, _ = opaque.derive_shared_keys(s_ikm, preamble) + server_mac = opaque.mac(s_km2, opaque.sha256_hash(preamble)) + + server_resp1 = login_pb2.LoginResponse( + opaque_response=login_pb2.OpaqueLoginResponse( + initial_response=login_pb2.InitialOpaqueLoginResponse( + server_nonce=server_login_nonce, + server_public_keyshare=server_ephemeral_pub, + server_mac=server_mac, + evaluated_message=evaluated_msg, + masking_nonce=masking_nonce, + masked_response=masked_response, + ) + ) + ) + + # 4. Client completes handshake + client_req2 = client_auth.final_request(server_resp1) + client_mac = client_req2.opaque_request.final_request.client_mac + + # 5. Server verifies client MAC + expected_client_mac = opaque.mac( + s_km3, + opaque.sha256_hash(opaque.concat(preamble, server_mac)), + ) + self.assertEqual(client_mac, expected_client_mac) + + def test_final_request_invalid_masked_response_length(self): + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + auth = opaque.UserAuthenticator("user", "pass", params) + auth.initial_request() + + resp = login_pb2.LoginResponse( + opaque_response=login_pb2.OpaqueLoginResponse( + initial_response=login_pb2.InitialOpaqueLoginResponse( + masked_response=b"invalid_len", + ) + ) + ) + with self.assertRaises(ValueError) as cm: + auth.final_request(resp) + self.assertIn("Invalid masked response length", str(cm.exception)) + + def test_clear(self): + opaque._clear(None) + b = bytearray(b"hello") + opaque._clear(b) + self.assertEqual(b, bytearray(5)) + + def test_point_add_identity_and_inversion(self): + pt = (opaque.GX, opaque.GY) + self.assertEqual(opaque.point_add(None, pt), pt) + self.assertEqual(opaque.point_add(pt, None), pt) + self.assertIsNone(opaque.point_add(None, None)) + + inv_pt = (opaque.GX, (opaque.P - opaque.GY) % opaque.P) + self.assertIsNone(opaque.point_add(pt, inv_pt)) + + def test_point_mul_infinity_and_zero(self): + self.assertIsNone(opaque.point_mul(None, 5)) + self.assertIsNone(opaque.point_mul(opaque.G, 0)) + + def test_marshal_compressed_none(self): + with self.assertRaises(ValueError) as cm: + opaque.marshal_compressed(None) + self.assertIn("Point at infinity cannot be compressed", str(cm.exception)) + + def test_unmarshal_compressed_errors(self): + # Invalid length + with self.assertRaises(ValueError) as cm: + opaque.unmarshal_compressed(b"\x02" * 32) + self.assertIn("Invalid compressed point length", str(cm.exception)) + + # Invalid prefix + with self.assertRaises(ValueError) as cm: + opaque.unmarshal_compressed(b"\x04" + b"\x00" * 32) + self.assertIn("Invalid compressed point prefix", str(cm.exception)) + + # x coordinate exceeds P + with self.assertRaises(ValueError) as cm: + opaque.unmarshal_compressed(b"\x02" + (opaque.P + 1).to_bytes(32, "big")) + self.assertIn("x coordinate exceeds field prime", str(cm.exception)) + + # Point not on curve (find an x that produces a quadratic non-residue) + for candidate_x in range(1, 100): + rhs = ( + pow(candidate_x, 3, opaque.P) + opaque.A * candidate_x + opaque.B + ) % opaque.P + if pow(rhs, opaque.P_MINUS_1_OVER_2, opaque.P) != 1: + with self.assertRaises(ValueError) as cm: + opaque.unmarshal_compressed( + b"\x02" + candidate_x.to_bytes(32, "big") + ) + self.assertIn("Point is not on curve", str(cm.exception)) + break + + def test_expand_message_xmd_oversize_dst(self): + res = opaque.expand_message_xmd(b"msg", b"D" * 256, 32) + self.assertEqual(len(res), 32) + + def test_map_to_curve_sswu_zero(self): + pt = opaque.map_to_curve_sswu(0) + self.assertIsInstance(pt, tuple) + self.assertEqual(len(pt), 2) + + def test_hash_to_curve_p256_infinity(self): + from unittest import mock + + with mock.patch( + "google.cloud.spanner_v1.omni.opaque.point_add", return_value=None + ): + with self.assertRaises(ValueError) as cm: + opaque.hash_to_curve_p256(b"msg", b"dst") + self.assertIn("Hash to curve produced point at infinity", str(cm.exception)) + + def test_validate_hash_parameters_errors(self): + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(None) + self.assertIn("hash_parameters cannot be None", str(cm.exception)) + + # Missing argon2_id_parameters on proto message + proto_empty = authentication_pb2.HashParameters() + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(proto_empty) + self.assertIn( + "hash_parameters must contain non-nil argon2_id_parameters", + str(cm.exception), + ) + + # Missing argon2_id_parameters on non-proto object + class EmptyParams: + argon2_id_parameters = None + + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(EmptyParams()) + self.assertIn( + "hash_parameters must contain non-nil argon2_id_parameters", + str(cm.exception), + ) + + # Invalid memory usage + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=7, + parallelism=4, + hash_size=32, + ) + ) + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(params) + self.assertIn("Invalid Argon2Id memory usage", str(cm.exception)) + + # Invalid parallelism + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=0, + hash_size=32, + ) + ) + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(params) + self.assertIn("Invalid Argon2Id parallelism", str(cm.exception)) + + # Invalid iteration count (0 and 11) + params_low = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=0, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(params_low) + self.assertIn("Invalid Argon2Id iteration count", str(cm.exception)) + + params_high = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=11, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(params_high) + self.assertIn("Invalid Argon2Id iteration count", str(cm.exception)) + + # Valid non-proto hash parameters + class ValidNonProtoParams: + argon2_id_parameters = authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + + opaque._validate_hash_parameters(ValidNonProtoParams()) + + # Invalid hash_size + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=0, + ) + ) + with self.assertRaises(ValueError) as cm: + opaque._validate_hash_parameters(params) + self.assertIn("Invalid Argon2Id hash size", str(cm.exception)) + + def test_random_oracle_sha256_large_max_val(self): + max_val = (1 << 264) + 1 + res = opaque.random_oracle_sha256(b"seed", max_val) + self.assertEqual(len(res), 34) + + def test_derive_key_pair_zero_priv(self): + from unittest import mock + + with mock.patch( + "google.cloud.spanner_v1.omni.opaque.random_oracle_sha256", + return_value=b"\x00" * 32, + ): + pub, priv = opaque.derive_key_pair(b"seed", b"info") + self.assertEqual(priv, (1).to_bytes(32, "big")) + self.assertEqual(pub, opaque.marshal_compressed(opaque.G)) + + def test_blind_and_finalize_edge_cases(self): + with self.assertRaises(ValueError) as cm: + opaque.blind(b"") + self.assertIn("Password cannot be empty", str(cm.exception)) + + # Explicit blind scalar + explicit_scalar = (5).to_bytes(32, "big") + blinded, returned_scalar = opaque.blind(b"pass", blind_scalar=explicit_scalar) + self.assertEqual(returned_scalar, explicit_scalar) + + with self.assertRaises(ValueError) as cm: + opaque.finalize(b"", b"evaluated") + self.assertIn("Blind scalar cannot be empty", str(cm.exception)) + + def test_recover_client_mismatched_auth_tag(self): + with self.assertRaises(ValueError) as cm: + opaque.recover_client( + "user", + b"\x01" * 32, + b"\x02" * 32, + b"\x00" * 32, + opaque.marshal_compressed(opaque.G), + ) + self.assertIn("Auth tag mismatch", str(cm.exception)) + + def test_final_request_edge_cases(self): + from unittest import mock + + params = authentication_pb2.HashParameters( + argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( + iteration_count=3, + memory_usage=64 * 1024, + parallelism=4, + hash_size=32, + ) + ) + auth = opaque.UserAuthenticator("user", "pass", params) + auth.initial_request() + + with self.assertRaises(ValueError) as cm: + auth.final_request(None) + self.assertIn("initial_response cannot be None", str(cm.exception)) + + empty_resp = login_pb2.LoginResponse() + with self.assertRaises(ValueError) as cm: + auth.final_request(empty_resp) + self.assertIn("Expected initial opaque response from server", str(cm.exception)) + + final_only_resp = login_pb2.LoginResponse() + final_only_resp.opaque_response.final_response.SetInParent() + with self.assertRaises(ValueError) as cm: + auth.final_request(final_only_resp) + self.assertIn("Expected initial opaque response from server", str(cm.exception)) + + # Test blind when secrets.randbelow initially returns 0 + with mock.patch("secrets.randbelow", side_effect=[0, 7]): + blinded_msg, blind_s = opaque.blind(b"testpass") + self.assertEqual(blind_s, (7).to_bytes(32, "big")) + + # Test server MAC mismatch + oprf_seed = opaque.nonce() + oprf_key_seed = opaque.expand(oprf_seed, b"userOprfKey", 32) + _, oprf_priv = opaque.derive_key_pair(oprf_key_seed, b"OPAQUE-DeriveKeyPair") + h_pt = opaque.hash_to_curve_p256(b"pass", opaque.LOGIN_DOMAIN_SEPARATION_TAG) + prf_pt = opaque.point_mul(h_pt, int.from_bytes(oprf_priv, "big")) + prf = opaque.marshal_compressed(prf_pt) + stretched_oprf = opaque.stretch(prf, params) + randomized_password = opaque.extract(opaque.concat(prf, stretched_oprf)) + + server_key_seed = opaque.nonce() + server_pub, _ = opaque.derive_key_pair( + server_key_seed, opaque.DIFFIE_HELLMAN_KEY_INFO + ) + envelope_nonce = opaque.nonce() + auth_key = opaque.expand( + randomized_password, envelope_nonce + opaque.AUTH_KEY_INFO, 32 + ) + auth_tag = opaque.mac(auth_key, envelope_nonce + server_pub + b"user") + serialized_envelope = opaque.concat(server_pub, envelope_nonce, auth_tag) + masking_key = opaque.expand(randomized_password, opaque.MASKING_KEY_INFO, 32) + masking_nonce = opaque.nonce() + credential_pad = opaque.expand( + masking_key, + opaque.concat(masking_nonce, b"CredentialResponsePad"), + len(serialized_envelope), + ) + masked_response = opaque.xor_bytes(serialized_envelope, credential_pad) + + blinded_pt = opaque.point_mul(h_pt, int.from_bytes(auth._blind, "big")) + eval_pt = opaque.point_mul(blinded_pt, int.from_bytes(oprf_priv, "big")) + server_ephemeral_pub, _ = opaque.derive_key_pair( + opaque.nonce(), opaque.DIFFIE_HELLMAN_KEY_INFO + ) + + resp = login_pb2.LoginResponse( + opaque_response=login_pb2.OpaqueLoginResponse( + initial_response=login_pb2.InitialOpaqueLoginResponse( + server_nonce=opaque.nonce(), + server_public_keyshare=server_ephemeral_pub, + server_mac=b"bad_server_mac" * 2 + b"\x00" * 4, + evaluated_message=opaque.marshal_compressed(eval_pt), + masking_nonce=masking_nonce, + masked_response=masked_response, + ) + ) + ) + + # Test server MAC mismatch on auth (which matches auth._blind) + with self.assertRaises(ValueError) as cm: + auth.final_request(resp) + self.assertIn("Server MAC mismatch", str(cm.exception)) + + # Test xor_bytes returning invalid envelope length + auth_short = opaque.UserAuthenticator("user", "pass", params) + auth_short.initial_request() + blinded_pt_short = opaque.point_mul( + h_pt, int.from_bytes(auth_short._blind, "big") + ) + eval_pt_short = opaque.point_mul( + blinded_pt_short, int.from_bytes(oprf_priv, "big") + ) + resp_short = login_pb2.LoginResponse( + opaque_response=login_pb2.OpaqueLoginResponse( + initial_response=login_pb2.InitialOpaqueLoginResponse( + server_nonce=opaque.nonce(), + server_public_keyshare=server_ephemeral_pub, + server_mac=b"bad_server_mac" * 2 + b"\x00" * 4, + evaluated_message=opaque.marshal_compressed(eval_pt_short), + masking_nonce=masking_nonce, + masked_response=masked_response, + ) + ) + ) + with mock.patch( + "google.cloud.spanner_v1.omni.opaque.xor_bytes", return_value=b"too_short" + ): + with self.assertRaises(ValueError) as cm: + auth_short.final_request(resp_short) + self.assertIn("Invalid serialized envelope length", str(cm.exception)) + + +if __name__ == "__main__": + unittest.main() diff --git a/packages/google-cloud-spanner/tests/unit/spanner_dbapi/test_connect.py b/packages/google-cloud-spanner/tests/unit/spanner_dbapi/test_connect.py index 4496884027f7..7ac0fc7aa1c9 100644 --- a/packages/google-cloud-spanner/tests/unit/spanner_dbapi/test_connect.py +++ b/packages/google-cloud-spanner/tests/unit/spanner_dbapi/test_connect.py @@ -183,3 +183,108 @@ def test_w_auto_partition_mode(self, mock_client): self.assertIsInstance(connection, Connection) self.assertTrue(connection.auto_partition_mode) + + def test_connect_omni_with_username_and_password(self, mock_client): + from google.cloud.spanner_dbapi import Connection, connect + from google.cloud.spanner_v1.omni.credentials import ( + SpannerOmniCredentials, + ) + + connection = connect( + INSTANCE, + DATABASE, + instance_type="omni", + client_options={"api_endpoint": "omni-host:15000"}, + username="test_user", + password="test_password", + ) + + self.assertIsInstance(connection, Connection) + mock_client.assert_called_once_with( + project="default", + credentials=mock.ANY, + client_info=mock.ANY, + route_to_leader_enabled=True, + client_options=mock.ANY, + use_plain_text=False, + ca_certificate=None, + client_certificate=None, + client_key=None, + instance_type="omni", + ) + creds = mock_client.call_args_list[0][1]["credentials"] + self.assertIsInstance(creds, SpannerOmniCredentials) + self.assertEqual(creds.username, "test_user") + + def test_connect_omni_partial_auth_raises_value_error(self, mock_client): + from google.cloud.spanner_dbapi import connect + + with self.assertRaises(ValueError) as ctx: + connect( + INSTANCE, + DATABASE, + instance_type="omni", + client_options={"api_endpoint": "omni-host:15000"}, + username="test_user", + ) + self.assertIn( + "Both username and password must be specified for Omni authentication", + str(ctx.exception), + ) + + with self.assertRaises(ValueError) as ctx: + connect( + INSTANCE, + DATABASE, + instance_type="omni", + client_options={"api_endpoint": "omni-host:15000"}, + password="test_password", + ) + self.assertIn( + "Both username and password must be specified for Omni authentication", + str(ctx.exception), + ) + + def test_connect_omni_client_options_variants(self, mock_client): + from google.api_core.client_options import ClientOptions + + from google.cloud.spanner_dbapi import connect + + # client_options is None + connect( + INSTANCE, + DATABASE, + instance_type="omni", + experimental_host="omni-host:15000", + client_options=None, + ) + + # client_options is dict + connect( + INSTANCE, + DATABASE, + instance_type="omni", + experimental_host="omni-host:15000", + client_options={"quota_project_id": "test-project"}, + ) + + # client_options is ClientOptions object with api_endpoint + opts_with_ep = ClientOptions(api_endpoint="omni-host:15000") + connect( + INSTANCE, + DATABASE, + instance_type="omni", + client_options=opts_with_ep, + ) + + # Missing host when instance_type='omni' raises ValueError + with self.assertRaises(ValueError) as ctx: + connect( + INSTANCE, + DATABASE, + instance_type="omni", + ) + self.assertIn( + "Host must be set for connecting to Spanner Omni instances", + str(ctx.exception), + ) diff --git a/packages/google-cloud-spanner/tests/unit/test__helpers.py b/packages/google-cloud-spanner/tests/unit/test__helpers.py index 0a6e9594b167..3776bbd26141 100644 --- a/packages/google-cloud-spanner/tests/unit/test__helpers.py +++ b/packages/google-cloud-spanner/tests/unit/test__helpers.py @@ -1862,3 +1862,190 @@ def test_large_values(self): self.assertEqual(result.months, case["expected_months"]) self.assertEqual(result.days, case["expected_days"]) self.assertEqual(result.nanos, case["expected_nanos"]) + + +class TestCreateSpannerOmniTransport(unittest.TestCase): + def test_create_spanner_omni_transport_plaintext_with_auth_interceptor(self): + import grpc + + from google.cloud.spanner_v1 import _helpers + + mock_factory = mock.MagicMock() + mock_creds = mock.MagicMock() + mock_interceptor = mock.MagicMock(spec=grpc.UnaryUnaryClientInterceptor) + mock_creds.create_auth_interceptor.return_value = mock_interceptor + + with mock.patch("grpc.insecure_channel") as mock_insecure: + with mock.patch("grpc.intercept_channel") as mock_intercept: + _helpers._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + credentials=mock_creds, + ) + mock_insecure.assert_called_once_with(target="localhost:9010") + mock_intercept.assert_called_once_with( + mock_insecure.return_value, mock_interceptor + ) + mock_factory.assert_called_once_with( + channel=mock_intercept.return_value, credentials=mock_creds + ) + + def test_create_spanner_omni_transport_tls_and_mtls(self): + from google.cloud.spanner_v1 import _helpers + + mock_factory = mock.MagicMock() + with mock.patch("builtins.open", mock.mock_open(read_data=b"cert_data")): + with mock.patch("grpc.ssl_channel_credentials") as mock_ssl_creds: + with mock.patch("grpc.secure_channel") as mock_secure: + # TLS only + _helpers._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate=None, + client_key=None, + ) + mock_ssl_creds.assert_called_with(root_certificates=b"cert_data") + mock_secure.assert_called_with( + "omni-host:15000", mock_ssl_creds.return_value + ) + + # mTLS + _helpers._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate="client.pem", + client_key="key.pem", + ) + mock_ssl_creds.assert_called_with( + root_certificates=b"cert_data", + private_key=b"cert_data", + certificate_chain=b"cert_data", + ) + + def test_create_spanner_omni_transport_validation_errors(self): + from google.cloud.spanner_v1 import _helpers + + mock_factory = mock.MagicMock() + # Missing ca_certificate + with self.assertRaises(ValueError) as cm: + _helpers._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate=None, + client_certificate=None, + client_key=None, + ) + self.assertIn("TLS/mTLS connection requires ca_certificate", str(cm.exception)) + + # Missing client_key when client_certificate provided + with mock.patch("builtins.open", mock.mock_open(read_data=b"cert_data")): + with self.assertRaises(ValueError) as cm: + _helpers._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate="client.pem", + client_key=None, + ) + self.assertIn( + "Both client_certificate and client_key must be provided for mTLS connection", + str(cm.exception), + ) + + # Missing client_certificate when client_key provided + with mock.patch("builtins.open", mock.mock_open(read_data=b"cert_data")): + with self.assertRaises(ValueError) as cm: + _helpers._create_spanner_omni_transport( + mock_factory, + "omni-host:15000", + use_plain_text=False, + ca_certificate="ca.pem", + client_certificate=None, + client_key="key.pem", + ) + self.assertIn( + "Both client_certificate and client_key must be provided for mTLS connection", + str(cm.exception), + ) + + def test_create_spanner_omni_transport_interceptors_and_credentials_fallback(self): + import grpc + from google.auth.credentials import AnonymousCredentials + + from google.cloud.spanner_v1 import _helpers + + mock_factory = mock.MagicMock() + existing_interceptor = mock.MagicMock(spec=grpc.UnaryUnaryClientInterceptor) + mock_creds = mock.MagicMock(spec=["create_auth_interceptor"]) + auth_interceptor = mock.MagicMock(spec=grpc.UnaryUnaryClientInterceptor) + mock_creds.create_auth_interceptor.return_value = auth_interceptor + + # Case 1: credentials with interceptor and existing interceptors + with mock.patch("grpc.insecure_channel") as mock_insecure: + with mock.patch("grpc.intercept_channel") as mock_intercept: + _helpers._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + interceptors=[existing_interceptor], + credentials=mock_creds, + ) + mock_intercept.assert_called_once_with( + mock_insecure.return_value, existing_interceptor, auth_interceptor + ) + mock_factory.assert_called_once_with( + channel=mock_intercept.return_value, credentials=mock_creds + ) + + # Case 2: credentials without create_auth_interceptor, no interceptors + mock_plain_creds = mock.MagicMock(spec=[]) + mock_factory.reset_mock() + with mock.patch("grpc.insecure_channel") as mock_insecure: + with mock.patch("grpc.intercept_channel") as mock_intercept: + _helpers._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + interceptors=None, + credentials=mock_plain_creds, + ) + mock_intercept.assert_not_called() + mock_factory.assert_called_once_with( + channel=mock_insecure.return_value, credentials=mock_plain_creds + ) + + # Case 3: credentials is None -> uses AnonymousCredentials + mock_factory.reset_mock() + with mock.patch("grpc.insecure_channel") as mock_insecure: + with mock.patch("grpc.intercept_channel") as mock_intercept: + _helpers._create_spanner_omni_transport( + mock_factory, + "localhost:9010", + use_plain_text=True, + ca_certificate=None, + client_certificate=None, + client_key=None, + interceptors=None, + credentials=None, + ) + mock_intercept.assert_not_called() + self.assertEqual(mock_factory.call_count, 1) + self.assertIsInstance( + mock_factory.call_args[1]["credentials"], AnonymousCredentials + ) diff --git a/packages/google-cloud-spanner/tests/unit/test_client.py b/packages/google-cloud-spanner/tests/unit/test_client.py index 69bb317f58e0..ac279e031230 100644 --- a/packages/google-cloud-spanner/tests/unit/test_client.py +++ b/packages/google-cloud-spanner/tests/unit/test_client.py @@ -988,3 +988,117 @@ def test_constructor_w_invalid_instance_type_raises_value_error(self): self.assertIn( "instance_type must be one of 'cloud' or 'omni'", str(ctx.exception) ) + + def test_constructor_w_omni_username_password(self): + from google.cloud.spanner_v1.client import InstanceType + from google.cloud.spanner_v1.omni.credentials import ( + SpannerOmniCredentials, + ) + + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + password="test_password", + ) + self.assertEqual(client.project, "default") + self.assertEqual(client.instance_type, InstanceType.OMNI) + self.assertEqual(client._host, "omni-host:15000") + self.assertIsInstance(client._credentials, SpannerOmniCredentials) + self.assertEqual(client._credentials.username, "test_user") + self.assertEqual(client._credentials.target, "omni-host:15000") + + def test_constructor_w_omni_partial_credentials_raises_value_error(self): + from google.cloud.spanner_v1.client import InstanceType + + with self.assertRaises(ValueError) as ctx: + self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + ) + self.assertIn( + "Both username and password must be specified for Omni authentication", + str(ctx.exception), + ) + + with self.assertRaises(ValueError) as ctx: + self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + password="test_password", + ) + self.assertIn( + "Both username and password must be specified for Omni authentication", + str(ctx.exception), + ) + + def test_constructor_w_username_password_on_cloud_raises_value_error(self): + creds = build_scoped_credentials() + with self.assertRaises(ValueError) as ctx: + self._make_one( + project=self.PROJECT, + credentials=creds, + username="test_user", + password="test_password", + ) + self.assertIn( + "username and password can only be used when instance_type='omni'.", + str(ctx.exception), + ) + + def test_instance_admin_api_omni(self): + from google.cloud.spanner_v1.client import InstanceType + + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + password="test_password", + use_plain_text=True, + ) + + inst_module = "google.cloud.spanner_v1.client.InstanceAdminClient" + with mock.patch(inst_module) as instance_admin_client: + api = client.instance_admin_api + self.assertIs(api, instance_admin_client.return_value) + instance_admin_client.assert_called_once() + called_kw = instance_admin_client.call_args[1] + self.assertIn("transport", called_kw) + + def test_database_admin_api_omni(self): + from google.cloud.spanner_v1.client import InstanceType + + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + username="test_user", + password="test_password", + use_plain_text=True, + ) + + db_module = "google.cloud.spanner_v1.client.DatabaseAdminClient" + with mock.patch(db_module) as database_admin_client: + api = client.database_admin_api + self.assertIs(api, database_admin_client.return_value) + database_admin_client.assert_called_once() + called_kw = database_admin_client.call_args[1] + self.assertIn("transport", called_kw) + + def test_constructor_w_omni_explicit_credentials_instance(self): + from google.cloud.spanner_v1.client import InstanceType + from google.cloud.spanner_v1.omni.credentials import SpannerOmniCredentials + + creds = SpannerOmniCredentials("user", "pass", "omni-host:15000") + client = self._make_one( + project=self.PROJECT, + client_options={"api_endpoint": "omni-host:15000"}, + instance_type=InstanceType.OMNI, + credentials=creds, + ) + self.assertIs(client._credentials, creds)