diff --git a/google/genai/_api_client.py b/google/genai/_api_client.py index 4d90ba86e..9cfb84cc5 100644 --- a/google/genai/_api_client.py +++ b/google/genai/_api_client.py @@ -832,6 +832,7 @@ def __init__( # Initialize the aiohttp client sessions. self._aiohttp_sessions: dict[Any, Any] = {} + self._async_client_session_request_args: dict[str, Any] = {} if self._use_aiohttp(): try: import aiohttp # pylint: disable=g-import-not-at-top @@ -1393,7 +1394,11 @@ 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 ( + self._authorized_session._is_mtls + and 'googleapis.com' in url + and self.location not in ['us', 'eu'] # mtls is not supported in multi-regions + ): if 'sandbox' in url: url = url.replace( 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' @@ -1471,7 +1476,11 @@ 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 ( + session._is_mtls # type: ignore[union-attr] + and 'googleapis.com' in url + and self.location not in ['us', 'eu'] # mtls is not supported in multi-regions + ): if 'sandbox' in url: url = url.replace( 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' @@ -1541,7 +1550,11 @@ 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 ( + session._is_mtls # type: ignore[union-attr] + and 'googleapis.com' in url + and self.location not in ['us', 'eu'] # mtls is not supported in multi-regions + ): if 'sandbox' in url: url = url.replace( 'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com' diff --git a/google/genai/tests/client/test_client_requests.py b/google/genai/tests/client/test_client_requests.py index 4bbc35492..e8d224abe 100644 --- a/google/genai/tests/client/test_client_requests.py +++ b/google/genai/tests/client/test_client_requests.py @@ -16,11 +16,33 @@ """Tests for client behavior when issuing requests.""" +from unittest import mock +import asyncio +import pytest + +try: + import aiohttp + + AIOHTTP_NOT_INSTALLED = False +except ImportError: + AIOHTTP_NOT_INSTALLED = True + aiohttp = mock.MagicMock() + from ... import _api_client as api_client from ... import Client from ... import types +requires_aiohttp = pytest.mark.skipif( + AIOHTTP_NOT_INSTALLED, reason='aiohttp is not installed, skipping test.' +) + + +@pytest.fixture(autouse=True) +def reset_has_aiohttp(): + api_client.has_aiohttp = not AIOHTTP_NOT_INSTALLED + + def build_test_client(monkeypatch): monkeypatch.setenv('GOOGLE_API_KEY', 'google_api_key') return Client() @@ -194,7 +216,6 @@ def test_build_request_with_resource_scope_with_project_and_location( assert request.url == 'https://custom-base-url.com/publishers/google/models/gemini-3-pro-preview' - def build_test_client_no_env_vars(monkeypatch): monkeypatch.delenv('GOOGLE_API_KEY', raising=False) monkeypatch.delenv('GEMINI_API_KEY', raising=False) @@ -218,4 +239,304 @@ def test_build_request_with_custom_base_url_no_env_vars(monkeypatch): 'test/path', {'key': 'value'}, ) - assert request.url == 'https://custom-base-url.com' \ No newline at end of file + assert request.url == 'https://custom-base-url.com' + + +def test_sync_request_mtls_regional(monkeypatch): + mock_session = mock.MagicMock() + mock_session._is_mtls = True + mock_response = mock.MagicMock() + mock_response.status_code = 200 + mock_response.headers = {} + mock_response.text = '{}' + mock_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us-central1') + client._api_client._authorized_session = mock_session + with mock.patch.object( + client._api_client, '_use_google_auth_sync', return_value=True + ), mock.patch.object( + client._api_client, '_access_token', return_value='mock_token' + ): + client._api_client._request_once( + api_client.HttpRequest( + method='GET', + url='https://us-central1-aiplatform.googleapis.com/v1beta1/models', + headers={}, + data={}, + ) + ) + assert ( + mock_session.request.call_args.kwargs['url'] + == 'https://us-central1-aiplatform.mtls.googleapis.com/v1beta1/models' + ) + + +def test_sync_request_mtls_sandbox(monkeypatch): + mock_session = mock.MagicMock() + mock_session._is_mtls = True + mock_response = mock.MagicMock() + mock_response.status_code = 200 + mock_response.headers = {} + mock_response.text = '{}' + mock_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us-central1') + client._api_client._authorized_session = mock_session + with mock.patch.object( + client._api_client, '_use_google_auth_sync', return_value=True + ), mock.patch.object( + client._api_client, '_access_token', return_value='mock_token' + ): + client._api_client._request_once( + api_client.HttpRequest( + method='GET', + url='https://us-central1-aiplatform.sandbox.googleapis.com/v1beta1/models', + headers={}, + data={}, + ) + ) + assert ( + mock_session.request.call_args.kwargs['url'] + == 'https://us-central1-aiplatform.mtls.sandbox.googleapis.com/v1beta1/models' + ) + + +def test_sync_request_mtls_multi_regional(monkeypatch): + mock_session = mock.MagicMock() + mock_session._is_mtls = True + mock_response = mock.MagicMock() + mock_response.status_code = 200 + mock_response.headers = {} + mock_response.text = '{}' + mock_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us') + client._api_client._authorized_session = mock_session + with mock.patch.object( + client._api_client, '_use_google_auth_sync', return_value=True + ), mock.patch.object( + client._api_client, '_access_token', return_value='mock_token' + ): + client._api_client._request_once( + api_client.HttpRequest( + method='GET', + url='https://aiplatform.us.rep.googleapis.com/v1beta1/models', + headers={}, + data={}, + ) + ) + assert ( + mock_session.request.call_args.kwargs['url'] + == 'https://aiplatform.us.rep.googleapis.com/v1beta1/models' + ) + + +@requires_aiohttp +@pytest.mark.asyncio +async def test_async_request_mtls_regional(monkeypatch): + mock_async_session = mock.AsyncMock() + mock_async_session._is_mtls = True + mock_async_session.closed = False + mock_response = mock.MagicMock(spec=aiohttp.ClientResponse) + mock_response.status = 200 + mock_response.headers = {} + mock_async_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us-central1') + client._api_client._http_options.aiohttp_client = mock_async_session + client._api_client._aiohttp_sessions[asyncio.get_running_loop()] = ( + mock_async_session + ) + client._api_client._async_client_session_request_args = {} + with mock.patch.object( + client._api_client, '_use_aiohttp', return_value=True + ), mock.patch.object( + client._api_client, '_use_google_auth_async', return_value=True + ), mock.patch.object( + client._api_client, '_async_access_token', return_value='mock_token' + ), mock.patch( + 'google.auth.transport.mtls.default_client_cert_source', + return_value='mock_cert_source', + ): + await client._api_client._async_request_once( + api_client.HttpRequest( + method='GET', + url='https://us-central1-aiplatform.googleapis.com/v1beta1/models', + headers={}, + data={}, + ) + ) + mock_async_session.configure_mtls_channel.assert_called_once_with( + 'mock_cert_source' + ) + assert ( + mock_async_session.request.call_args.kwargs['url'] + == 'https://us-central1-aiplatform.mtls.googleapis.com/v1beta1/models' + ) + + +@requires_aiohttp +@pytest.mark.asyncio +async def test_async_request_mtls_sandbox(monkeypatch): + mock_async_session = mock.AsyncMock() + mock_async_session._is_mtls = True + mock_async_session.closed = False + mock_response = mock.MagicMock(spec=aiohttp.ClientResponse) + mock_response.status = 200 + mock_response.headers = {} + mock_async_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us-central1') + client._api_client._http_options.aiohttp_client = mock_async_session + client._api_client._aiohttp_sessions[asyncio.get_running_loop()] = ( + mock_async_session + ) + client._api_client._async_client_session_request_args = {} + with mock.patch.object( + client._api_client, '_use_aiohttp', return_value=True + ), mock.patch.object( + client._api_client, '_use_google_auth_async', return_value=True + ), mock.patch.object( + client._api_client, '_async_access_token', return_value='mock_token' + ), mock.patch( + 'google.auth.transport.mtls.default_client_cert_source', + return_value='mock_cert_source', + ): + await client._api_client._async_request_once( + api_client.HttpRequest( + method='GET', + url='https://us-central1-aiplatform.sandbox.googleapis.com/v1beta1/models', + headers={}, + data={}, + ) + ) + assert ( + mock_async_session.request.call_args.kwargs['url'] + == 'https://us-central1-aiplatform.mtls.sandbox.googleapis.com/v1beta1/models' + ) + + +@requires_aiohttp +@pytest.mark.asyncio +async def test_async_request_mtls_multi_regional(monkeypatch): + mock_async_session = mock.AsyncMock() + mock_async_session._is_mtls = True + mock_async_session.closed = False + mock_response = mock.MagicMock(spec=aiohttp.ClientResponse) + mock_response.status = 200 + mock_response.headers = {} + mock_async_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us') + client._api_client._http_options.aiohttp_client = mock_async_session + client._api_client._aiohttp_sessions[asyncio.get_running_loop()] = ( + mock_async_session + ) + client._api_client._async_client_session_request_args = {} + with mock.patch.object( + client._api_client, '_use_aiohttp', return_value=True + ), mock.patch.object( + client._api_client, '_use_google_auth_async', return_value=True + ), mock.patch.object( + client._api_client, '_async_access_token', return_value='mock_token' + ), mock.patch( + 'google.auth.transport.mtls.default_client_cert_source', + return_value='mock_cert_source', + ): + await client._api_client._async_request_once( + api_client.HttpRequest( + method='GET', + url='https://aiplatform.us.rep.googleapis.com/v1beta1/models', + headers={}, + data={}, + ) + ) + assert ( + mock_async_session.request.call_args.kwargs['url'] + == 'https://aiplatform.us.rep.googleapis.com/v1beta1/models' + ) + + +@requires_aiohttp +@pytest.mark.asyncio +async def test_async_stream_request_mtls_regional(monkeypatch): + mock_async_session = mock.AsyncMock() + mock_async_session._is_mtls = True + mock_async_session.closed = False + mock_response = mock.MagicMock(spec=aiohttp.ClientResponse) + mock_response.status = 200 + mock_response.headers = {} + mock_async_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us-central1') + client._api_client._http_options.aiohttp_client = mock_async_session + client._api_client._aiohttp_sessions[asyncio.get_running_loop()] = ( + mock_async_session + ) + client._api_client._async_client_session_request_args = {} + with mock.patch.object( + client._api_client, '_use_aiohttp', return_value=True + ), mock.patch.object( + client._api_client, '_use_google_auth_async', return_value=True + ), mock.patch.object( + client._api_client, '_async_access_token', return_value='mock_token' + ), mock.patch( + 'google.auth.transport.mtls.default_client_cert_source', + return_value='mock_cert_source', + ): + await client._api_client._async_request_once( + api_client.HttpRequest( + method='GET', + url='https://us-central1-aiplatform.googleapis.com/v1beta1/models', + headers={}, + data={}, + ), + stream=True, + ) + assert ( + mock_async_session.request.call_args.kwargs['url'] + == 'https://us-central1-aiplatform.mtls.googleapis.com/v1beta1/models' + ) + + +@requires_aiohttp +@pytest.mark.asyncio +async def test_async_stream_request_mtls_multi_regional(monkeypatch): + mock_async_session = mock.AsyncMock() + mock_async_session._is_mtls = True + mock_async_session.closed = False + mock_response = mock.MagicMock(spec=aiohttp.ClientResponse) + mock_response.status = 200 + mock_response.headers = {} + mock_async_session.request.return_value = mock_response + + client = Client(vertexai=True, project='test-project', location='us') + client._api_client._http_options.aiohttp_client = mock_async_session + client._api_client._aiohttp_sessions[asyncio.get_running_loop()] = ( + mock_async_session + ) + client._api_client._async_client_session_request_args = {} + with mock.patch.object( + client._api_client, '_use_aiohttp', return_value=True + ), mock.patch.object( + client._api_client, '_use_google_auth_async', return_value=True + ), mock.patch.object( + client._api_client, '_async_access_token', return_value='mock_token' + ), mock.patch( + 'google.auth.transport.mtls.default_client_cert_source', + return_value='mock_cert_source', + ): + await client._api_client._async_request_once( + api_client.HttpRequest( + method='GET', + url='https://aiplatform.us.rep.googleapis.com/v1beta1/models', + headers={}, + data={}, + ), + stream=True, + ) + assert ( + mock_async_session.request.call_args.kwargs['url'] + == 'https://aiplatform.us.rep.googleapis.com/v1beta1/models' + )