Skip to content

Commit 3466586

Browse files
xuanyang15copybara-github
authored andcommitted
feat: Add mTLS support to Google API tools
Co-authored-by: Xuan Yang <xygoogle@google.com> PiperOrigin-RevId: 941350837
1 parent f706a1e commit 3466586

8 files changed

Lines changed: 325 additions & 20 deletions

File tree

src/google/adk/integrations/parameter_manager/parameter_client.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@
2525
from google.oauth2 import service_account
2626

2727
from ... import version
28-
from ...utils import _mtls_utils
28+
from ...utils._mtls_utils import get_api_endpoint
2929

3030
USER_AGENT = f"google-adk/{version.__version__}"
3131

@@ -112,7 +112,7 @@ def __init__(
112112
client_options = None
113113
if location:
114114
client_options = {
115-
"api_endpoint": _mtls_utils.get_api_endpoint(
115+
"api_endpoint": get_api_endpoint(
116116
location,
117117
_DEFAULT_REGIONAL_ENDPOINT_TEMPLATE,
118118
_DEFAULT_MTLS_REGIONAL_ENDPOINT_TEMPLATE,

src/google/adk/tools/google_api_tool/google_api_toolset.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,22 +14,28 @@
1414

1515
from __future__ import annotations
1616

17+
import logging
1718
from typing import Dict
1819
from typing import List
1920
from typing import Optional
2021
from typing import Union
2122

23+
import httpx
2224
from typing_extensions import override
2325

2426
from ...agents.readonly_context import ReadonlyContext
2527
from ...auth.auth_credential import ServiceAccount
2628
from ...auth.auth_schemes import OpenIdConnectWithConfig
2729
from ...tools.base_toolset import BaseToolset
2830
from ...tools.base_toolset import ToolPredicate
31+
from ...utils._mtls_utils import MtlsClientCerts
32+
from ...utils._mtls_utils import use_client_cert_effective
2933
from ..openapi_tool import OpenAPIToolset
3034
from .google_api_tool import GoogleApiTool
3135
from .googleapi_to_openapi_converter import GoogleApiToOpenApiConverter
3236

37+
logger = logging.getLogger('google_adk.' + __name__)
38+
3339

3440
class GoogleApiToolset(BaseToolset):
3541
"""Google API Toolset contains tools for interacting with Google APIs.
@@ -75,6 +81,20 @@ def __init__(
7581
self._additional_headers = additional_headers
7682
self._additional_scopes = additional_scopes
7783
self._discovery_url = discovery_url
84+
85+
self._httpx_client_factory = None
86+
use_client_cert = use_client_cert_effective()
87+
88+
if use_client_cert:
89+
self._mtls_certs = MtlsClientCerts()
90+
cert_path, key_path, passphrase = self._mtls_certs.get_certs()
91+
if cert_path and key_path and passphrase:
92+
93+
def client_factory():
94+
return httpx.AsyncClient(cert=(cert_path, key_path, passphrase))
95+
96+
self._httpx_client_factory = client_factory
97+
7898
self._openapi_toolset = self._load_toolset_with_oidc_auth()
7999

80100
@override
@@ -134,6 +154,7 @@ def _load_toolset_with_oidc_auth(self) -> OpenAPIToolset:
134154
grant_types_supported=['authorization_code'],
135155
scopes=scopes,
136156
),
157+
httpx_client_factory=self._httpx_client_factory,
137158
)
138159

139160
def configure_auth(self, client_id: str, client_secret: str):
@@ -147,3 +168,5 @@ def configure_sa_auth(self, service_account: ServiceAccount):
147168
async def close(self):
148169
if self._openapi_toolset:
149170
await self._openapi_toolset.close()
171+
if hasattr(self, '_mtls_certs') and self._mtls_certs:
172+
self._mtls_certs.close()

src/google/adk/tools/google_api_tool/googleapi_to_openapi_converter.py

Lines changed: 44 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -17,13 +17,18 @@
1717
import argparse
1818
import json
1919
import logging
20+
import socket
2021
from typing import Any
2122
from typing import Dict
2223
from typing import List
2324

2425
# Google API client
2526
from googleapiclient.discovery import build
2627
from googleapiclient.errors import HttpError
28+
import httplib2
29+
30+
from ...utils._mtls_utils import MtlsClientCerts
31+
from ...utils._mtls_utils import use_client_cert_effective
2732

2833
# Configure logging
2934
logger = logging.getLogger("google_adk." + __name__)
@@ -54,6 +59,8 @@ def __init__(
5459
"paths": {},
5560
"components": {"schemas": {}, "securitySchemes": {}},
5661
}
62+
self._use_client_cert = use_client_cert_effective()
63+
self._mtls_certs = MtlsClientCerts() if self._use_client_cert else None
5764

5865
def fetch_google_api_spec(self) -> None:
5966
"""Fetches the Google API specification using discovery service."""
@@ -63,15 +70,42 @@ def fetch_google_api_spec(self) -> None:
6370
self._api_name,
6471
self._api_version,
6572
)
73+
74+
# Determine if we should use mTLS
75+
# self._use_client_cert is already initialized in __init__
76+
77+
http_client = None
78+
discovery_url = self._discovery_url
79+
80+
if self._use_client_cert and self._mtls_certs:
81+
cert_path, key_path, passphrase = self._mtls_certs.get_certs()
82+
if cert_path and key_path and passphrase:
83+
# Set default HTTP timeout similar to googleapiclient.http.build_http()
84+
http_timeout = socket.getdefaulttimeout() or 60
85+
http_client = httplib2.Http(timeout=http_timeout)
86+
try:
87+
http_client.redirect_codes = http_client.redirect_codes - {308}
88+
except AttributeError:
89+
pass
90+
http_client.add_certificate(key_path, cert_path, "", passphrase)
91+
92+
if not discovery_url:
93+
discovery_url = "https://www.mtls.googleapis.com/discovery/v1/apis/{api}/{apiVersion}/rest"
94+
6695
# Build a resource object for the specified API
67-
if self._discovery_url:
96+
if discovery_url:
6897
self._google_api_resource = build(
6998
self._api_name,
7099
self._api_version,
71-
discoveryServiceUrl=self._discovery_url,
100+
discoveryServiceUrl=discovery_url,
101+
http=http_client,
72102
)
73103
else:
74-
self._google_api_resource = build(self._api_name, self._api_version)
104+
self._google_api_resource = build(
105+
self._api_name,
106+
self._api_version,
107+
http=http_client,
108+
)
75109

76110
# Access the underlying API discovery document
77111
self._google_api_spec = self._google_api_resource._rootDesc
@@ -136,9 +170,13 @@ def _convert_info(self) -> None:
136170

137171
def _convert_servers(self) -> None:
138172
"""Convert server information."""
139-
base_url = self._google_api_spec.get(
140-
"rootUrl", ""
141-
) + self._google_api_spec.get("servicePath", "")
173+
use_client_cert = getattr(self, "_use_client_cert", False)
174+
if use_client_cert and "mtlsRootUrl" in self._google_api_spec:
175+
root_url = self._google_api_spec["mtlsRootUrl"]
176+
else:
177+
root_url = self._google_api_spec.get("rootUrl", "")
178+
179+
base_url = root_url + self._google_api_spec.get("servicePath", "")
142180

143181
# Remove trailing slash if present
144182
if base_url.endswith("/"):

src/google/adk/utils/_mtls_utils.py

Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,8 @@
1919
import enum
2020
import logging
2121
import os
22+
import tempfile
23+
import threading
2224
from typing import TYPE_CHECKING
2325
from urllib.parse import urlsplit
2426
from urllib.parse import urlunsplit
@@ -149,3 +151,69 @@ def configure_session_for_mtls(session: requests.Session) -> bool:
149151
if is_mtls:
150152
session.mount("https://", _MutualTlsAdapter(cert, key))
151153
return bool(is_mtls)
154+
155+
156+
class MtlsClientCerts:
157+
"""Manages the creation and lifecycle of client certificates for mTLS.
158+
159+
Extracts certificates to a temporary directory that is automatically cleaned up
160+
when the instance is garbage collected.
161+
"""
162+
163+
def __init__(self) -> None:
164+
self._tempdir: tempfile.TemporaryDirectory[str] | None = None
165+
self.cert_path: str | None = None
166+
self.key_path: str | None = None
167+
self.passphrase: bytes | None = None
168+
self._lock = threading.Lock()
169+
self._initialized = False
170+
171+
def get_certs(self) -> tuple[str | None, str | None, bytes | None]:
172+
"""Extracts and returns the certificate paths and passphrase.
173+
174+
Returns:
175+
A tuple of (cert_path, key_path, passphrase) if client certificates
176+
are available, otherwise (None, None, None).
177+
"""
178+
with self._lock:
179+
if self._initialized:
180+
return self.cert_path, self.key_path, self.passphrase
181+
182+
if not mtls.has_default_client_cert_source():
183+
self._initialized = True
184+
return None, None, None
185+
186+
self._tempdir = tempfile.TemporaryDirectory()
187+
cert_path_tmp = os.path.join(self._tempdir.name, "cert.pem")
188+
key_path_tmp = os.path.join(self._tempdir.name, "key.pem")
189+
190+
try:
191+
cert_source = mtls.default_client_encrypted_cert_source(
192+
cert_path_tmp, key_path_tmp
193+
)
194+
_, _, passphrase = cert_source()
195+
except Exception as e:
196+
# If extraction fails, we should fail loud.
197+
self._tempdir.cleanup()
198+
self._tempdir = None
199+
raise RuntimeError(
200+
f"Failed to extract default client certificates for mTLS: {e}"
201+
) from e
202+
203+
self.cert_path = cert_path_tmp
204+
self.key_path = key_path_tmp
205+
self.passphrase = passphrase
206+
self._initialized = True
207+
208+
return self.cert_path, self.key_path, self.passphrase
209+
210+
def close(self) -> None:
211+
"""Manually cleans up the temporary directory."""
212+
with self._lock:
213+
if self._tempdir:
214+
self._tempdir.cleanup()
215+
self._tempdir = None
216+
self.cert_path = None
217+
self.key_path = None
218+
self.passphrase = None
219+
self._initialized = False

tests/unittests/integrations/parameter_manager/test_parameter_client.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -115,7 +115,7 @@ def test_init_with_auth_token(self, mock_pm_client_class):
115115
"google.adk.integrations.parameter_manager.parameter_client.default_service_credential"
116116
)
117117
@patch(
118-
"google.adk.integrations.parameter_manager.parameter_client._mtls_utils.get_api_endpoint"
118+
"google.adk.integrations.parameter_manager.parameter_client.get_api_endpoint"
119119
)
120120
def test_init_with_location(
121121
self,

tests/unittests/tools/google_api_tool/test_google_api_toolset.py

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -523,3 +523,44 @@ def test_init_with_tool_name_prefix(
523523
)
524524

525525
assert tool_set.tool_name_prefix == tool_name_prefix
526+
527+
@mock.patch(
528+
"google.adk.tools.google_api_tool.google_api_toolset.OpenAPIToolset"
529+
)
530+
@mock.patch(
531+
"google.adk.tools.google_api_tool.google_api_toolset.GoogleApiToOpenApiConverter"
532+
)
533+
@mock.patch(
534+
"google.adk.tools.google_api_tool.google_api_toolset.MtlsClientCerts"
535+
)
536+
@mock.patch(
537+
"google.adk.tools.google_api_tool.google_api_toolset.use_client_cert_effective"
538+
)
539+
async def test_mtls_cleanup_on_close(
540+
self,
541+
mock_use_client_cert,
542+
mock_mtls_certs_class,
543+
mock_converter_class,
544+
mock_openapi_toolset_class,
545+
):
546+
"""Test that mTLS temp files are cleaned up on close."""
547+
mock_converter_class.return_value = mock.MagicMock()
548+
mock_openapi_toolset_instance = mock.MagicMock()
549+
mock_openapi_toolset_instance.close = mock.AsyncMock()
550+
mock_openapi_toolset_class.return_value = mock_openapi_toolset_instance
551+
552+
mock_use_client_cert.return_value = True
553+
mock_mtls_certs_instance = mock.MagicMock()
554+
mock_mtls_certs_instance.get_certs.return_value = ("cert", "key", b"pass")
555+
mock_mtls_certs_class.return_value = mock_mtls_certs_instance
556+
557+
tool_set = GoogleApiToolset(
558+
api_name=TEST_API_NAME, api_version=TEST_API_VERSION
559+
)
560+
561+
assert tool_set._httpx_client_factory is not None
562+
563+
await tool_set.close()
564+
565+
mock_openapi_toolset_instance.close.assert_called_once()
566+
mock_mtls_certs_instance.close.assert_called_once()

tests/unittests/tools/google_api_tool/test_googleapi_to_openapi_converter.py

Lines changed: 52 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -197,6 +197,14 @@ def calendar_api_spec():
197197
}
198198

199199

200+
@pytest.fixture(autouse=True)
201+
def disable_mtls_by_default(monkeypatch):
202+
monkeypatch.setattr(
203+
"google.auth.transport.mtls.should_use_client_cert",
204+
lambda: False,
205+
)
206+
207+
200208
@pytest.fixture
201209
def converter():
202210
"""Fixture that provides a basic converter instance."""
@@ -279,7 +287,50 @@ def test_fetch_google_api_spec_with_discovery_url(
279287

280288
assert converter._google_api_spec == calendar_api_spec
281289
mock_build.assert_called_once_with(
282-
"calendar", "v3", discoveryServiceUrl=discovery_url
290+
"calendar", "v3", discoveryServiceUrl=discovery_url, http=None
291+
)
292+
293+
def test_fetch_google_api_spec_with_mtls(
294+
self, monkeypatch, mock_api_resource, calendar_api_spec
295+
):
296+
"""Test fetching Google API specification with mTLS enabled."""
297+
mock_build = MagicMock(return_value=mock_api_resource)
298+
monkeypatch.setattr(
299+
"google.adk.tools.google_api_tool.googleapi_to_openapi_converter.build",
300+
mock_build,
301+
)
302+
303+
# Enable mTLS
304+
monkeypatch.setattr(
305+
"google.auth.transport.mtls.should_use_client_cert",
306+
lambda: True,
307+
)
308+
monkeypatch.setattr(
309+
"google.auth.transport.mtls.has_default_client_cert_source",
310+
lambda: True,
311+
)
312+
313+
mock_cert_source = MagicMock(
314+
return_value=("/path/to/cert", "/path/to/key", b"passphrase")
315+
)
316+
monkeypatch.setattr(
317+
"google.auth.transport.mtls.default_client_encrypted_cert_source",
318+
lambda c, k: mock_cert_source,
319+
)
320+
321+
converter = GoogleApiToOpenApiConverter("calendar", "v3")
322+
converter.fetch_google_api_spec()
323+
324+
assert converter._google_api_spec == calendar_api_spec
325+
326+
# Verify build was called with the http parameter set and mtls url
327+
mock_build.assert_called_once()
328+
_, kwargs = mock_build.call_args
329+
assert "http" in kwargs
330+
assert kwargs["http"] is not None
331+
assert (
332+
kwargs["discoveryServiceUrl"]
333+
== "https://www.mtls.googleapis.com/discovery/v1/apis/{api}/{apiVersion}/rest"
283334
)
284335

285336
def test_fetch_google_api_spec_error(self, monkeypatch, converter):

0 commit comments

Comments
 (0)