mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: Make OpenAPI tool async
Merge https://github.com/google/adk-python/pull/2872 Closes https://github.com/google/adk-python/issues/787 The OpenAPI tool has been ported to the httpx client to make requests truly asynchronous. COPYBARA_INTEGRATE_REVIEW=https://github.com/google/adk-python/pull/2872 from condorcet:async_openapi_tool bf83f73af93f624126462fb0bd41fef27c53a0b6 PiperOrigin-RevId: 864250822
This commit is contained in:
committed by
Copybara-Service
parent
2770012cec
commit
9290b96626
@@ -45,6 +45,7 @@ dependencies = [
|
||||
"google-cloud-storage>=2.18.0, <4.0.0", # For GCS Artifact service
|
||||
"google-genai>=1.56.0, <2.0.0", # Google GenAI SDK
|
||||
"graphviz>=0.20.2, <1.0.0", # Graphviz for graph rendering
|
||||
"httpx>=0.27.0, <1.0.0", # HTTP client library
|
||||
"jsonschema>=4.23.0, <5.0.0", # Agent Builder config validation
|
||||
"mcp>=1.23.0, <2.0.0", # For MCP Toolset
|
||||
"opentelemetry-api>=1.37.0, <=1.37.0", # OpenTelemetry - limit upper version for sdk and api to not risk breaking changes from unstable _logs package.
|
||||
|
||||
@@ -28,9 +28,9 @@ from fastapi.openapi.models import HTTPBearer
|
||||
from fastapi.openapi.models import OAuth2
|
||||
from fastapi.openapi.models import OpenIdConnect
|
||||
from fastapi.openapi.models import Schema
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from pydantic import ValidationError
|
||||
import requests
|
||||
|
||||
from ....auth.auth_credential import AuthCredential
|
||||
from ....auth.auth_credential import AuthCredentialTypes
|
||||
@@ -289,14 +289,14 @@ def openid_url_to_scheme_credential(
|
||||
Raises:
|
||||
ValueError: If the OpenID URL is invalid, fetching fails, or required
|
||||
fields are missing.
|
||||
requests.exceptions.RequestException: If there's an error during the
|
||||
httpx.HTTPStatusError or httpx.RequestError: If there's an error during the
|
||||
HTTP request.
|
||||
"""
|
||||
try:
|
||||
response = requests.get(openid_url, timeout=10)
|
||||
response = httpx.get(openid_url, timeout=10)
|
||||
response.raise_for_status()
|
||||
config_dict = response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
except httpx.RequestError as e:
|
||||
raise ValueError(
|
||||
f"Failed to fetch OpenID configuration from {openid_url}: {e}"
|
||||
) from e
|
||||
|
||||
@@ -28,7 +28,7 @@ from typing import Union
|
||||
from fastapi.openapi.models import Operation
|
||||
from fastapi.openapi.models import Schema
|
||||
from google.genai.types import FunctionDeclaration
|
||||
import requests
|
||||
import httpx
|
||||
from typing_extensions import override
|
||||
|
||||
from ....agents.readonly_context import ReadonlyContext
|
||||
@@ -312,7 +312,7 @@ class RestApiTool(BaseTool):
|
||||
|
||||
Returns:
|
||||
A dictionary containing the request parameters for the API call. This
|
||||
initializes a requests.request() call.
|
||||
initializes an httpx.AsyncClient.request() call.
|
||||
|
||||
Example:
|
||||
self._prepare_request_params({"input_id": "test-id"})
|
||||
@@ -497,17 +497,7 @@ class RestApiTool(BaseTool):
|
||||
if provider_headers:
|
||||
request_params.setdefault("headers", {}).update(provider_headers)
|
||||
|
||||
# Log the API request
|
||||
self._logger.debug(
|
||||
"API Request: %s %s",
|
||||
request_params.get("method", "").upper(),
|
||||
request_params.get("url", ""),
|
||||
)
|
||||
self._logger.debug("API Request params: %s", request_params.get("params"))
|
||||
if "json" in request_params:
|
||||
self._logger.debug("API Request body: %s", request_params.get("json"))
|
||||
|
||||
response = requests.request(**request_params)
|
||||
response = await _request(**request_params)
|
||||
|
||||
# Log the API response
|
||||
self._logger.debug(
|
||||
@@ -519,11 +509,9 @@ class RestApiTool(BaseTool):
|
||||
|
||||
# Parse API response
|
||||
try:
|
||||
response.raise_for_status() # Raise HTTPError for bad responses
|
||||
result = response.json() # Try to decode JSON
|
||||
self._logger.debug("API Response body: %s", result)
|
||||
return result
|
||||
except requests.exceptions.HTTPError:
|
||||
response.raise_for_status() # Raise HTTPStatusError for bad responses
|
||||
return response.json() # Try to decode JSON
|
||||
except httpx.HTTPStatusError:
|
||||
error_details = response.content.decode("utf-8")
|
||||
self._logger.warning(
|
||||
"API call failed for tool %s: Status %d - %s",
|
||||
@@ -556,3 +544,10 @@ class RestApiTool(BaseTool):
|
||||
f' auth_scheme="{self.auth_scheme}",'
|
||||
f' auth_credential="{self.auth_credential}")'
|
||||
)
|
||||
|
||||
|
||||
async def _request(**request_params) -> httpx.Response:
|
||||
async with httpx.AsyncClient(
|
||||
verify=request_params.pop("verify", True)
|
||||
) as client:
|
||||
return await client.request(**request_params)
|
||||
|
||||
@@ -36,8 +36,8 @@ from google.adk.tools.openapi_tool.auth.auth_helpers import openid_url_to_scheme
|
||||
from google.adk.tools.openapi_tool.auth.auth_helpers import service_account_dict_to_scheme_credential
|
||||
from google.adk.tools.openapi_tool.auth.auth_helpers import service_account_scheme_credential
|
||||
from google.adk.tools.openapi_tool.auth.auth_helpers import token_to_scheme_credential
|
||||
import httpx
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
|
||||
def test_token_to_scheme_credential_api_key_header():
|
||||
@@ -272,7 +272,7 @@ def test_openid_dict_to_scheme_credential_missing_credential_fields():
|
||||
openid_dict_to_scheme_credential(config_dict, scopes, credential_dict)
|
||||
|
||||
|
||||
@patch("requests.get")
|
||||
@patch("httpx.get")
|
||||
def test_openid_url_to_scheme_credential(mock_get):
|
||||
mock_response = {
|
||||
"authorization_endpoint": "auth_url",
|
||||
@@ -303,7 +303,7 @@ def test_openid_url_to_scheme_credential(mock_get):
|
||||
mock_get.assert_called_once_with("openid_url", timeout=10)
|
||||
|
||||
|
||||
@patch("requests.get")
|
||||
@patch("httpx.get")
|
||||
def test_openid_url_to_scheme_credential_no_openid_url(mock_get):
|
||||
mock_response = {
|
||||
"authorization_endpoint": "auth_url",
|
||||
@@ -326,9 +326,9 @@ def test_openid_url_to_scheme_credential_no_openid_url(mock_get):
|
||||
assert scheme.openIdConnectUrl == "openid_url"
|
||||
|
||||
|
||||
@patch("requests.get")
|
||||
@patch("httpx.get")
|
||||
def test_openid_url_to_scheme_credential_request_exception(mock_get):
|
||||
mock_get.side_effect = requests.exceptions.RequestException("Test Error")
|
||||
mock_get.side_effect = httpx.RequestError("Test Error", request=None)
|
||||
credential_dict = {"client_id": "client_id", "client_secret": "client_secret"}
|
||||
|
||||
with pytest.raises(
|
||||
@@ -337,7 +337,7 @@ def test_openid_url_to_scheme_credential_request_exception(mock_get):
|
||||
openid_url_to_scheme_credential("openid_url", [], credential_dict)
|
||||
|
||||
|
||||
@patch("requests.get")
|
||||
@patch("httpx.get")
|
||||
def test_openid_url_to_scheme_credential_invalid_json(mock_get):
|
||||
mock_get.return_value.json.side_effect = ValueError("Invalid JSON")
|
||||
mock_get.return_value.raise_for_status.return_value = None
|
||||
|
||||
@@ -41,6 +41,7 @@ from google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool import snak
|
||||
from google.adk.tools.tool_context import ToolContext
|
||||
from google.genai.types import FunctionDeclaration
|
||||
from google.genai.types import Schema
|
||||
import httpx
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
@@ -246,7 +247,7 @@ class TestRestApiTool:
|
||||
}
|
||||
|
||||
@patch(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.requests.request"
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool._request"
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_success(
|
||||
@@ -278,7 +279,7 @@ class TestRestApiTool:
|
||||
assert result == {"result": "success"}
|
||||
|
||||
@patch(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.requests.request"
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool._request"
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_http_failure(
|
||||
@@ -293,8 +294,15 @@ class TestRestApiTool:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 500
|
||||
mock_response.content = b"Internal Server Error"
|
||||
mock_response.raise_for_status.side_effect = requests.exceptions.HTTPError(
|
||||
"500 Server Error"
|
||||
|
||||
# Create a proper HTTPStatusError with request and response
|
||||
mock_http_request = MagicMock(spec=httpx.Request)
|
||||
mock_response.raise_for_status = MagicMock(
|
||||
side_effect=httpx.HTTPStatusError(
|
||||
"500 Server Error",
|
||||
request=mock_http_request,
|
||||
response=mock_response,
|
||||
)
|
||||
)
|
||||
mock_request.return_value = mock_response
|
||||
|
||||
@@ -321,7 +329,7 @@ class TestRestApiTool:
|
||||
}
|
||||
|
||||
@patch(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.requests.request"
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool._request"
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_auth_pending(
|
||||
@@ -359,7 +367,7 @@ class TestRestApiTool:
|
||||
}
|
||||
|
||||
@patch(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.requests.request"
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool._request"
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_required_param_defaults(
|
||||
@@ -1052,6 +1060,11 @@ class TestRestApiTool:
|
||||
mock_response.json.return_value = {"result": "success"}
|
||||
mock_response.configure_mock(status_code=200)
|
||||
|
||||
mock_client = mock.create_autospec(
|
||||
httpx.AsyncClient, instance=True, spec_set=True
|
||||
)
|
||||
mock_client.request = AsyncMock(return_value=mock_response)
|
||||
|
||||
tool = RestApiTool(
|
||||
name="test_tool",
|
||||
description="Test Tool",
|
||||
@@ -1063,14 +1076,15 @@ class TestRestApiTool:
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
requests, "request", return_value=mock_response, autospec=True
|
||||
httpx, "AsyncClient", return_value=mock_client, autospec=True
|
||||
) as mock_request:
|
||||
|
||||
await tool.call(args={}, tool_context=mock_tool_context)
|
||||
|
||||
assert mock_request.called
|
||||
_, call_kwargs = mock_request.call_args
|
||||
if expected_verify_in_call is None:
|
||||
assert "verify" not in call_kwargs
|
||||
assert "verify" not in call_kwargs or call_kwargs["verify"] is True
|
||||
else:
|
||||
assert call_kwargs["verify"] == expected_verify_in_call
|
||||
|
||||
@@ -1087,6 +1101,11 @@ class TestRestApiTool:
|
||||
mock_response.json.return_value = {"result": "success"}
|
||||
mock_response.configure_mock(status_code=200)
|
||||
|
||||
mock_client = mock.create_autospec(
|
||||
httpx.AsyncClient, instance=True, spec_set=True
|
||||
)
|
||||
mock_client.request = AsyncMock(return_value=mock_response)
|
||||
|
||||
tool = RestApiTool(
|
||||
name="test_tool",
|
||||
description="Test Tool",
|
||||
@@ -1100,7 +1119,7 @@ class TestRestApiTool:
|
||||
tool.configure_ssl_verify(ca_bundle_path)
|
||||
|
||||
with patch.object(
|
||||
requests, "request", return_value=mock_response
|
||||
httpx, "AsyncClient", return_value=mock_client, autospec=True
|
||||
) as mock_request:
|
||||
await tool.call(args={}, tool_context=mock_tool_context)
|
||||
|
||||
@@ -1169,13 +1188,14 @@ class TestRestApiTool:
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
requests, "request", return_value=mock_response, autospec=True
|
||||
httpx.AsyncClient, "request", return_value=mock_response, autospec=True
|
||||
) as mock_request:
|
||||
await tool.call(args={}, tool_context=mock_tool_context)
|
||||
|
||||
# Verify the headers were added to the request
|
||||
assert mock_request.called
|
||||
_, call_kwargs = mock_request.call_args
|
||||
|
||||
assert call_kwargs["headers"]["X-Custom-Header"] == "custom-value"
|
||||
assert call_kwargs["headers"]["X-Request-ID"] == "12345"
|
||||
|
||||
@@ -1210,7 +1230,7 @@ class TestRestApiTool:
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
requests, "request", return_value=mock_response, autospec=True
|
||||
httpx.AsyncClient, "request", return_value=mock_response, autospec=True
|
||||
):
|
||||
await tool.call(args={}, tool_context=mock_tool_context)
|
||||
|
||||
@@ -1242,7 +1262,7 @@ class TestRestApiTool:
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
requests, "request", return_value=mock_response, autospec=True
|
||||
httpx.AsyncClient, "request", return_value=mock_response, autospec=True
|
||||
):
|
||||
result = await tool.call(args={}, tool_context=mock_tool_context)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user