From 97d5821da9591645919fbcac5179c1327882fc7c Mon Sep 17 00:00:00 2001 From: Amy Wu Date: Wed, 7 Oct 2026 17:03:45 -0700 Subject: [PATCH] fix: Use the mTLS endpoint on the httpx and websocket transports PiperOrigin-RevId: 995448624 --- google/genai/_api_client.py | 143 +++++++++----- .../client/test_client_initialization.py | 179 ++++++++++++++++++ 2 files changed, 279 insertions(+), 43 deletions(-) diff --git a/google/genai/_api_client.py b/google/genai/_api_client.py index 541fc8038..b93830070 100644 --- a/google/genai/_api_client.py +++ b/google/genai/_api_client.py @@ -222,6 +222,25 @@ def join_url_path(base_url: str, path: str) -> str: return urlunparse(parsed_base._replace(path=base_path + '/' + path)) +def to_mtls_url(url: str) -> str: + """Rewrites a googleapis.com URL to its mTLS endpoint. + + For example, `aiplatform.googleapis.com` becomes + `aiplatform.mtls.googleapis.com` and `foo.sandbox.googleapis.com` becomes + `foo.mtls.sandbox.googleapis.com`. URLs on other hosts, and URLs that already + point at an mTLS endpoint, are returned unchanged. + """ + parsed = urlparse(url) + netloc = parsed.netloc + if netloc.endswith(('.mtls.googleapis.com', '.mtls.sandbox.googleapis.com')): + return url + for domain in ('sandbox.googleapis.com', 'googleapis.com'): + if netloc.endswith('.' + domain): + netloc = netloc[: -len(domain)] + 'mtls.' + domain + return urlunparse(parsed._replace(netloc=netloc)) + return url + + def load_auth(*, project: Union[str, None]) -> Tuple[Credentials, str]: """Loads google auth credentials and project id.""" credentials, loaded_project_id = google.auth.default( # type: ignore[no-untyped-call] @@ -898,6 +917,26 @@ def __init__( self._http_options, vertexai=bool(self.vertexai), ) + # The httpx and websocket transports present the default client certificate + # only through the SSL context the SDK creates, so when the caller supplies + # their own client or SSL context, only GOOGLE_API_USE_MTLS_ENDPOINT=always + # switches them to the mTLS endpoint. + client_args = self._http_options.client_args or {} + async_client_args = self._http_options.async_client_args or {} + custom_httpx_verify = bool( + client_args.get('verify') or async_client_args.get('verify') + ) + self._httpx_use_mtls_endpoint = self._use_mtls_endpoint( + sdk_ssl_ctx=not (self._http_options.httpx_client or custom_httpx_verify) + ) + self._async_httpx_use_mtls_endpoint = self._use_mtls_endpoint( + sdk_ssl_ctx=not ( + self._http_options.httpx_async_client or custom_httpx_verify + ) + ) + self._websocket_use_mtls_endpoint = self._use_mtls_endpoint( + sdk_ssl_ctx=not async_client_args.get('ssl') + ) self._retry = tenacity.Retrying(**retry_kwargs) self._async_retry = tenacity.AsyncRetrying(**retry_kwargs) @@ -913,6 +952,37 @@ def _use_google_auth_sync(self) -> bool: ) ) + def _use_mtls_endpoint(self, sdk_ssl_ctx: bool) -> bool: + """Returns whether a transport should send requests to the mTLS endpoint. + + Certificate-bound access tokens are only accepted on `mtls.googleapis.com`, + so a transport that presents the client certificate must also switch to + the mTLS endpoint. + + Args: + sdk_ssl_ctx: Whether the transport uses the SSL context created by the + SDK, which carries the default client certificate when one is + configured. + """ + if not self.vertexai: + return False + client_cert_available = bool( + sdk_ssl_ctx + and hasattr(mtls, 'should_use_client_cert') + and mtls.should_use_client_cert() # type: ignore[no-untyped-call] + and mtls.has_default_client_cert_source() # type: ignore[no-untyped-call] + ) + should_use_mtls_endpoint = getattr(mtls, 'should_use_mtls_endpoint', None) + if should_use_mtls_endpoint is None: + return client_cert_available + try: + return bool( + should_use_mtls_endpoint(client_cert_available=client_cert_available) + ) + except auth_exceptions.MutualTLSChannelError as e: + logger.warning('Failed to determine whether to use mTLS endpoint: %s', e) + return client_cert_available + def _use_google_auth_async(self) -> bool: try: import importlib @@ -1294,6 +1364,15 @@ def _use_aiohttp(self) -> bool: and (self._http_options.httpx_async_client is None) ) + def _httpx_url(self, url: str, *, is_async: bool) -> str: + """Returns the URL to send a request to with the httpx client.""" + use_mtls_endpoint = ( + self._async_httpx_use_mtls_endpoint + if is_async + else self._httpx_use_mtls_endpoint + ) + return to_mtls_url(url) if use_mtls_endpoint else url + def _websocket_base_url(self) -> str: has_sufficient_auth = (self.project and self.location) or self.api_key if self.custom_base_url and not has_sufficient_auth: @@ -1301,7 +1380,10 @@ def _websocket_base_url(self) -> str: # Enable custom url if auth is not sufficient. return self.custom_base_url url_parts = urlparse(self._http_options.base_url) - return url_parts._replace(scheme='wss').geturl() # type: ignore[arg-type, return-value] + url = url_parts._replace(scheme='wss').geturl() # type: ignore[arg-type] + if self._websocket_use_mtls_endpoint: + url = to_mtls_url(url) # type: ignore[arg-type] + return url # type: ignore[return-value] def _access_token(self) -> str: """Retrieves the access token for the credentials.""" @@ -1494,13 +1576,8 @@ def _request_once( self._authorized_session.configure_mtls_channel( client_cert_source ) # type: ignore[no-untyped-call] - if self._authorized_session._is_mtls and 'googleapis.com' in url: - if 'sandbox' in url: - url = url.replace( - 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' - ) - else: - url = url.replace('googleapis.com', 'mtls.googleapis.com') + if self._authorized_session._is_mtls: + url = to_mtls_url(url) response = self._authorized_session.request( # type: ignore[no-untyped-call] method=http_request.method.upper(), url=url, @@ -1512,7 +1589,7 @@ def _request_once( else: httpx_request = self._httpx_client.build_request( # type: ignore[union-attr] method=http_request.method, - url=http_request.url, + url=self._httpx_url(http_request.url, is_async=False), content=data, headers=http_request.headers, timeout=http_request.timeout, @@ -1572,13 +1649,8 @@ async def _async_request_once( await session.configure_mtls_channel( # type: ignore[union-attr] client_cert_source ) - if session._is_mtls and 'googleapis.com' in url: # type: ignore[union-attr] - if 'sandbox' in url: - url = url.replace( - 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' - ) - else: - url = url.replace('googleapis.com', 'mtls.googleapis.com') + if session._is_mtls: # type: ignore[union-attr] + url = to_mtls_url(url) try: response = await session.request( # type: ignore[union-attr] method=http_request.method, @@ -1625,7 +1697,7 @@ async def _async_request_once( # aiohttp is not available. Fall back to httpx. httpx_request = self._async_httpx_client.build_request( # type: ignore[union-attr] method=http_request.method, - url=http_request.url, + url=self._httpx_url(http_request.url, is_async=True), content=data, headers=http_request.headers, timeout=http_request.timeout, @@ -1645,13 +1717,8 @@ async def _async_request_once( await session.configure_mtls_channel( # type: ignore[union-attr] client_cert_source ) - if session._is_mtls and 'googleapis.com' in url: # type: ignore[union-attr] - if 'sandbox' in url: - url = url.replace( - 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' - ) - else: - url = url.replace('googleapis.com', 'mtls.googleapis.com') + if session._is_mtls: # type: ignore[union-attr] + url = to_mtls_url(url) try: response = await session.request( # type: ignore[union-attr] method=http_request.method, @@ -1708,7 +1775,7 @@ async def _async_request_once( # aiohttp is not available. Fall back to httpx. client_response = await self._async_httpx_client.request( # type: ignore[union-attr] method=http_request.method, - url=http_request.url, + url=self._httpx_url(http_request.url, is_async=True), headers=http_request.headers, content=data, timeout=http_request.timeout, @@ -2054,13 +2121,8 @@ def _write_chunks(chunks: Iterator[bytes]) -> None: self._authorized_session.configure_mtls_channel( client_cert_source ) # type: ignore[no-untyped-call] - if self._authorized_session._is_mtls and 'googleapis.com' in url: - if 'sandbox' in url: - url = url.replace( - 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' - ) - else: - url = url.replace('googleapis.com', 'mtls.googleapis.com') + if self._authorized_session._is_mtls: + url = to_mtls_url(url) if destination is not None: response = self._authorized_session.request( # type: ignore[no-untyped-call] method=http_request.method.upper(), @@ -2090,7 +2152,7 @@ def _write_chunks(chunks: Iterator[bytes]) -> None: if destination is not None: httpx_request = self._httpx_client.build_request( # type: ignore[union-attr] method=http_request.method, - url=http_request.url, + url=self._httpx_url(http_request.url, is_async=False), content=data, headers=http_request.headers, timeout=http_request.timeout, @@ -2105,7 +2167,7 @@ def _write_chunks(chunks: Iterator[bytes]) -> None: else: response = self._httpx_client.request( # type: ignore[union-attr] method=http_request.method, - url=http_request.url, + url=self._httpx_url(http_request.url, is_async=False), content=data, headers=http_request.headers, timeout=http_request.timeout, @@ -2428,13 +2490,8 @@ async def _write_chunks(chunks: AsyncIterator[bytes]) -> None: await session.configure_mtls_channel( # type: ignore[union-attr] client_cert_source ) - if session._is_mtls and 'googleapis.com' in url: # type: ignore[union-attr] - if 'sandbox' in url: - url = url.replace( - 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' - ) - else: - url = url.replace('googleapis.com', 'mtls.googleapis.com') + if session._is_mtls: # type: ignore[union-attr] + url = to_mtls_url(url) response = await session.request( # type: ignore[union-attr] method=http_request.method, url=url, @@ -2465,7 +2522,7 @@ async def _write_chunks(chunks: AsyncIterator[bytes]) -> None: if destination is not None: httpx_request = self._async_httpx_client.build_request( # type: ignore[union-attr] method=http_request.method, - url=http_request.url, + url=self._httpx_url(http_request.url, is_async=True), content=data, headers=http_request.headers, timeout=http_request.timeout, @@ -2485,7 +2542,7 @@ async def _write_chunks(chunks: AsyncIterator[bytes]) -> None: else: client_response = await self._async_httpx_client.request( # type: ignore[union-attr] method=http_request.method, - url=http_request.url, + url=self._httpx_url(http_request.url, is_async=True), headers=http_request.headers, content=data, timeout=http_request.timeout, diff --git a/google/genai/tests/client/test_client_initialization.py b/google/genai/tests/client/test_client_initialization.py index bb127855f..7a06848dc 100644 --- a/google/genai/tests/client/test_client_initialization.py +++ b/google/genai/tests/client/test_client_initialization.py @@ -2065,3 +2065,182 @@ async def run(): thread.join() assert len({id(session) for session in sessions}) == 3 + + +@pytest.mark.parametrize( + "url, expected", + [ + ( + "https://us-central1-aiplatform.googleapis.com/v1/models", + "https://us-central1-aiplatform.mtls.googleapis.com/v1/models", + ), + ( + "wss://aiplatform.googleapis.com/ws", + "wss://aiplatform.mtls.googleapis.com/ws", + ), + ( + "https://foo.sandbox.googleapis.com/v1", + "https://foo.mtls.sandbox.googleapis.com/v1", + ), + ( + "https://aiplatform.mtls.googleapis.com/v1", + "https://aiplatform.mtls.googleapis.com/v1", + ), + ( + "https://foo.mtls.sandbox.googleapis.com/v1", + "https://foo.mtls.sandbox.googleapis.com/v1", + ), + ( + "https://example.com/sandbox/googleapis.com", + "https://example.com/sandbox/googleapis.com", + ), + ], +) +def test_to_mtls_url(url, expected): + assert api_client.to_mtls_url(url) == expected + + +@pytest.fixture +def mock_client_cert(monkeypatch): + """Simulates an environment with a default client certificate configured.""" + from google.auth.transport import mtls + + monkeypatch.delenv("GOOGLE_API_USE_MTLS_ENDPOINT", raising=False) + monkeypatch.setattr( + mtls, "should_use_client_cert", lambda: True, raising=False + ) + monkeypatch.setattr(mtls, "has_default_client_cert_source", lambda: True) + monkeypatch.setattr(mtls, "get_default_ssl_context", lambda: None) + monkeypatch.setattr( + api_client.BaseApiClient, "_access_token", lambda self: "token" + ) + + async def _async_access_token(self): + return "token" + + monkeypatch.setattr( + api_client.BaseApiClient, "_async_access_token", _async_access_token + ) + + +def _mtls_test_client(**http_options): + return Client( + vertexai=True, + project="fake-project", + location="us-central1", + http_options=http_options, + ) + + +def test_sync_httpx_uses_mtls_endpoint_with_client_cert(mock_client_cert): + # Custom client_args without an SSL context keep the SDK's SSL context, but + # route requests through httpx rather than AuthorizedSession. + client = _mtls_test_client(client_args={"follow_redirects": True}) + assert not client._api_client._use_google_auth_sync() + mock_send = mock.Mock(return_value=httpx.Response(200, text="{}")) + + with mock.patch.object(api_client.SyncHttpxClient, "send", mock_send): + client._api_client.request("post", "models/gemini:generateContent", {}) + + assert mock_send.call_args[0][0].url.host == ( + "us-central1-aiplatform.mtls.googleapis.com" + ) + + +@pytest.mark.asyncio +async def test_async_httpx_uses_mtls_endpoint_with_client_cert( + mock_client_cert, +): + api_client.has_aiohttp = False + client = _mtls_test_client() + mock_request = mock.AsyncMock(return_value=httpx.Response(200, text="{}")) + + with mock.patch.object(api_client.AsyncHttpxClient, "request", mock_request): + await client._api_client.async_request( + "post", "models/gemini:generateContent", {} + ) + + assert mock_request.call_args.kwargs["url"].startswith( + "https://us-central1-aiplatform.mtls.googleapis.com/" + ) + + +@pytest.mark.asyncio +async def test_async_httpx_download_uses_mtls_endpoint_with_client_cert( + mock_client_cert, +): + api_client.has_aiohttp = False + client = _mtls_test_client() + mock_request = mock.AsyncMock(return_value=httpx.Response(200, content=b"")) + + with mock.patch.object(api_client.AsyncHttpxClient, "request", mock_request): + await client._api_client.async_download_file("files/abc:download") + + assert mock_request.call_args.kwargs["url"].startswith( + "https://us-central1-aiplatform.mtls.googleapis.com/" + ) + + +def test_websocket_uses_mtls_endpoint_with_client_cert(mock_client_cert): + client = _mtls_test_client() + assert client._api_client._websocket_base_url() == ( + "wss://us-central1-aiplatform.mtls.googleapis.com/" + ) + + +def test_mtls_endpoint_not_used_without_client_cert(monkeypatch): + from google.auth.transport import mtls + + monkeypatch.delenv("GOOGLE_API_USE_MTLS_ENDPOINT", raising=False) + monkeypatch.setattr( + mtls, "should_use_client_cert", lambda: False, raising=False + ) + client = _mtls_test_client() + assert client._api_client._websocket_base_url() == ( + "wss://us-central1-aiplatform.googleapis.com/" + ) + assert ( + client._api_client._httpx_url( + "https://us-central1-aiplatform.googleapis.com/v1", is_async=True + ) + == "https://us-central1-aiplatform.googleapis.com/v1" + ) + + +def test_mtls_endpoint_not_used_with_custom_transport(mock_client_cert): + client = _mtls_test_client( + httpx_async_client=httpx.AsyncClient(), + async_client_args={"ssl": ssl.create_default_context()}, + ) + url = "https://us-central1-aiplatform.googleapis.com/v1" + assert client._api_client._httpx_url(url, is_async=True) == url + assert client._api_client._websocket_base_url() == ( + "wss://us-central1-aiplatform.googleapis.com/" + ) + + +@pytest.mark.parametrize( + "use_mtls_endpoint, expected_host", + [ + ("always", "us-central1-aiplatform.mtls.googleapis.com"), + ("never", "us-central1-aiplatform.googleapis.com"), + ], +) +def test_mtls_endpoint_env_override( + mock_client_cert, monkeypatch, use_mtls_endpoint, expected_host +): + monkeypatch.setenv("GOOGLE_API_USE_MTLS_ENDPOINT", use_mtls_endpoint) + client = _mtls_test_client(httpx_async_client=httpx.AsyncClient()) + assert ( + client._api_client._httpx_url( + "https://us-central1-aiplatform.googleapis.com/v1", is_async=True + ) + == f"https://{expected_host}/v1" + ) + + +def test_mtls_endpoint_not_used_for_gemini_api(mock_client_cert): + client = Client(api_key="test-api-key") + assert client._api_client._websocket_base_url() == ( + "wss://generativelanguage.googleapis.com/" + )