Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
143 changes: 100 additions & 43 deletions google/genai/_api_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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)

Expand All @@ -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
Expand Down Expand Up @@ -1294,14 +1364,26 @@ 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:
# API gateway proxy can use the auth in custom headers, not url.
# 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."""
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
Loading
Loading