feat: allow setting agent/application name for BigQuery tools

This will allow tracking of tool usage per agent/application.

PiperOrigin-RevId: 800607186
This commit is contained in:
Google Team Member
2025-08-28 14:10:01 -07:00
committed by Copybara-Service
parent f4a8df0ba2
commit 11a2ffe35a
9 changed files with 259 additions and 33 deletions
@@ -15,9 +15,9 @@
from __future__ import annotations
import os
import re
from unittest import mock
import google.adk
from google.adk.tools.bigquery.client import get_bigquery_client
from google.auth.exceptions import DefaultCredentialsError
from google.oauth2.credentials import Credentials
@@ -109,8 +109,8 @@ def test_bigquery_client_project_set_with_env():
assert client.project == "test-gcp-project"
def test_bigquery_client_user_agent():
"""Test BigQuery client user agent."""
def test_bigquery_client_user_agent_default():
"""Test BigQuery client default user agent."""
with mock.patch(
"google.cloud.bigquery.client.Connection", autospec=True
) as mock_connection:
@@ -123,7 +123,33 @@ def test_bigquery_client_user_agent():
# Verify that the tracking user agent was set
client_info_arg = mock_connection.call_args[1].get("client_info")
assert client_info_arg is not None
assert re.search(
r"adk-bigquery-tool google-adk/([0-9A-Za-z._\-+/]+)",
client_info_arg.user_agent,
expected_user_agents = {
"adk-bigquery-tool",
f"google-adk/{google.adk.__version__}",
}
actual_user_agents = set(client_info_arg.user_agent.split())
assert expected_user_agents.issubset(actual_user_agents)
def test_bigquery_client_user_agent_custom():
"""Test BigQuery client custom user agent."""
with mock.patch(
"google.cloud.bigquery.client.Connection", autospec=True
) as mock_connection:
# Trigger the BigQuery client creation
get_bigquery_client(
project="test-gcp-project",
credentials=mock.create_autospec(Credentials, instance=True),
user_agent="custom_user_agent",
)
# Verify that the tracking user agent was set
client_info_arg = mock_connection.call_args[1].get("client_info")
assert client_info_arg is not None
expected_user_agents = {
"adk-bigquery-tool",
f"google-adk/{google.adk.__version__}",
"custom_user_agent",
}
actual_user_agents = set(client_info_arg.user_agent.split())
assert expected_user_agents.issubset(actual_user_agents)
@@ -18,19 +18,22 @@ import os
from unittest import mock
from google.adk.tools.bigquery import metadata_tool
from google.adk.tools.bigquery.config import BigQueryToolConfig
from google.auth.exceptions import DefaultCredentialsError
from google.cloud import bigquery
from google.oauth2.credentials import Credentials
import pytest
@mock.patch.dict(os.environ, {}, clear=True)
@mock.patch("google.cloud.bigquery.Client.list_datasets", autospec=True)
@mock.patch("google.auth.default", autospec=True)
def test_list_dataset_ids(mock_default_auth, mock_list_datasets):
"""Test list_dataset_ids tool invocation."""
def test_list_dataset_ids_no_default_auth(
mock_default_auth, mock_list_datasets
):
"""Test list_dataset_ids tool invocation involves no default auth."""
project = "my_project_id"
mock_credentials = mock.create_autospec(Credentials, instance=True)
tool_settings = BigQueryToolConfig()
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
@@ -42,7 +45,9 @@ def test_list_dataset_ids(mock_default_auth, mock_list_datasets):
bigquery.DatasetReference(project, "dataset1"),
bigquery.DatasetReference(project, "dataset2"),
]
result = metadata_tool.list_dataset_ids(project, mock_credentials)
result = metadata_tool.list_dataset_ids(
project, mock_credentials, tool_settings
)
assert result == ["dataset1", "dataset2"]
mock_default_auth.assert_not_called()
@@ -50,9 +55,10 @@ def test_list_dataset_ids(mock_default_auth, mock_list_datasets):
@mock.patch.dict(os.environ, {}, clear=True)
@mock.patch("google.cloud.bigquery.Client.get_dataset", autospec=True)
@mock.patch("google.auth.default", autospec=True)
def test_get_dataset_info(mock_default_auth, mock_get_dataset):
"""Test get_dataset_info tool invocation."""
def test_get_dataset_info_no_default_auth(mock_default_auth, mock_get_dataset):
"""Test get_dataset_info tool invocation involves no default auth."""
mock_credentials = mock.create_autospec(Credentials, instance=True)
tool_settings = BigQueryToolConfig()
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
@@ -64,7 +70,7 @@ def test_get_dataset_info(mock_default_auth, mock_get_dataset):
Credentials, instance=True
)
result = metadata_tool.get_dataset_info(
"my_project_id", "my_dataset_id", mock_credentials
"my_project_id", "my_dataset_id", mock_credentials, tool_settings
)
assert result != {
"status": "ERROR",
@@ -76,12 +82,13 @@ def test_get_dataset_info(mock_default_auth, mock_get_dataset):
@mock.patch.dict(os.environ, {}, clear=True)
@mock.patch("google.cloud.bigquery.Client.list_tables", autospec=True)
@mock.patch("google.auth.default", autospec=True)
def test_list_table_ids(mock_default_auth, mock_list_tables):
"""Test list_table_ids tool invocation."""
def test_list_table_ids_no_default_auth(mock_default_auth, mock_list_tables):
"""Test list_table_ids tool invocation involves no default auth."""
project = "my_project_id"
dataset = "my_dataset_id"
dataset_ref = bigquery.DatasetReference(project, dataset)
mock_credentials = mock.create_autospec(Credentials, instance=True)
tool_settings = BigQueryToolConfig()
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
@@ -93,7 +100,9 @@ def test_list_table_ids(mock_default_auth, mock_list_tables):
bigquery.TableReference(dataset_ref, "table1"),
bigquery.TableReference(dataset_ref, "table2"),
]
result = metadata_tool.list_table_ids(project, dataset, mock_credentials)
result = metadata_tool.list_table_ids(
project, dataset, mock_credentials, tool_settings
)
assert result == ["table1", "table2"]
mock_default_auth.assert_not_called()
@@ -101,9 +110,10 @@ def test_list_table_ids(mock_default_auth, mock_list_tables):
@mock.patch.dict(os.environ, {}, clear=True)
@mock.patch("google.cloud.bigquery.Client.get_table", autospec=True)
@mock.patch("google.auth.default", autospec=True)
def test_get_table_info(mock_default_auth, mock_get_table):
"""Test get_table_info tool invocation."""
def test_get_table_info_no_default_auth(mock_default_auth, mock_get_table):
"""Test get_table_info tool invocation involves no default auth."""
mock_credentials = mock.create_autospec(Credentials, instance=True)
tool_settings = BigQueryToolConfig()
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
@@ -113,10 +123,116 @@ def test_get_table_info(mock_default_auth, mock_get_table):
mock_get_table.return_value = mock.create_autospec(Credentials, instance=True)
result = metadata_tool.get_table_info(
"my_project_id", "my_dataset_id", "my_table_id", mock_credentials
"my_project_id",
"my_dataset_id",
"my_table_id",
mock_credentials,
tool_settings,
)
assert result != {
"status": "ERROR",
"error_details": "Your default credentials were not found",
}
mock_default_auth.assert_not_called()
@mock.patch(
"google.adk.tools.bigquery.client.get_bigquery_client", autospec=True
)
def test_list_dataset_ids_bq_client_creation(mock_get_bigquery_client):
"""Test BigQuery client creation params during list_dataset_ids tool invocation."""
bq_project = "my_project_id"
bq_credentials = mock.create_autospec(Credentials, instance=True)
application_name = "my-agent"
tool_settings = BigQueryToolConfig(application_name=application_name)
metadata_tool.list_dataset_ids(bq_project, bq_credentials, tool_settings)
mock_get_bigquery_client.assert_called_once()
assert len(mock_get_bigquery_client.call_args.kwargs) == 3
assert mock_get_bigquery_client.call_args.kwargs["project"] == bq_project
assert (
mock_get_bigquery_client.call_args.kwargs["credentials"] == bq_credentials
)
assert (
mock_get_bigquery_client.call_args.kwargs["user_agent"]
== application_name
)
@mock.patch(
"google.adk.tools.bigquery.client.get_bigquery_client", autospec=True
)
def test_get_dataset_info_bq_client_creation(mock_get_bigquery_client):
"""Test BigQuery client creation params during get_dataset_info tool invocation."""
bq_project = "my_project_id"
bq_dataset = "my_dataset_id"
bq_credentials = mock.create_autospec(Credentials, instance=True)
application_name = "my-agent"
tool_settings = BigQueryToolConfig(application_name=application_name)
metadata_tool.get_dataset_info(
bq_project, bq_dataset, bq_credentials, tool_settings
)
mock_get_bigquery_client.assert_called_once()
assert len(mock_get_bigquery_client.call_args.kwargs) == 3
assert mock_get_bigquery_client.call_args.kwargs["project"] == bq_project
assert (
mock_get_bigquery_client.call_args.kwargs["credentials"] == bq_credentials
)
assert (
mock_get_bigquery_client.call_args.kwargs["user_agent"]
== application_name
)
@mock.patch(
"google.adk.tools.bigquery.client.get_bigquery_client", autospec=True
)
def test_list_table_ids_bq_client_creation(mock_get_bigquery_client):
"""Test BigQuery client creation params during list_table_ids tool invocation."""
bq_project = "my_project_id"
bq_dataset = "my_dataset_id"
bq_credentials = mock.create_autospec(Credentials, instance=True)
application_name = "my-agent"
tool_settings = BigQueryToolConfig(application_name=application_name)
metadata_tool.list_table_ids(
bq_project, bq_dataset, bq_credentials, tool_settings
)
mock_get_bigquery_client.assert_called_once()
assert len(mock_get_bigquery_client.call_args.kwargs) == 3
assert mock_get_bigquery_client.call_args.kwargs["project"] == bq_project
assert (
mock_get_bigquery_client.call_args.kwargs["credentials"] == bq_credentials
)
assert (
mock_get_bigquery_client.call_args.kwargs["user_agent"]
== application_name
)
@mock.patch(
"google.adk.tools.bigquery.client.get_bigquery_client", autospec=True
)
def test_get_table_info_bq_client_creation(mock_get_bigquery_client):
"""Test BigQuery client creation params during get_table_info tool invocation."""
bq_project = "my_project_id"
bq_dataset = "my_dataset_id"
bq_table = "my_table_id"
bq_credentials = mock.create_autospec(Credentials, instance=True)
application_name = "my-agent"
tool_settings = BigQueryToolConfig(application_name=application_name)
metadata_tool.get_table_info(
bq_project, bq_dataset, bq_table, bq_credentials, tool_settings
)
mock_get_bigquery_client.assert_called_once()
assert len(mock_get_bigquery_client.call_args.kwargs) == 3
assert mock_get_bigquery_client.call_args.kwargs["project"] == bq_project
assert (
mock_get_bigquery_client.call_args.kwargs["credentials"] == bq_credentials
)
assert (
mock_get_bigquery_client.call_args.kwargs["user_agent"]
== application_name
)
@@ -983,3 +983,26 @@ def test_execute_sql_result_dtype(
# Test the tool worked without invoking default auth
result = execute_sql(project, query, credentials, tool_settings, tool_context)
assert result == {"status": "SUCCESS", "rows": tool_result_rows}
@mock.patch(
"google.adk.tools.bigquery.client.get_bigquery_client", autospec=True
)
def test_execute_sql_bq_client_creation(mock_get_bigquery_client):
"""Test BigQuery client creation params during execute_sql tool invocation."""
project = "my_project_id"
query = "SELECT 1"
credentials = mock.create_autospec(Credentials, instance=True)
application_name = "my-agent"
tool_settings = BigQueryToolConfig(application_name=application_name)
tool_context = mock.create_autospec(ToolContext, instance=True)
execute_sql(project, query, credentials, tool_settings, tool_context)
mock_get_bigquery_client.assert_called_once()
assert len(mock_get_bigquery_client.call_args.kwargs) == 3
assert mock_get_bigquery_client.call_args.kwargs["project"] == project
assert mock_get_bigquery_client.call_args.kwargs["credentials"] == credentials
assert (
mock_get_bigquery_client.call_args.kwargs["user_agent"]
== application_name
)
@@ -25,3 +25,12 @@ def test_bigquery_tool_config_experimental_warning():
match="Config defaults may have breaking change in the future.",
):
BigQueryToolConfig()
def test_bigquery_tool_config_invalid_application_name():
"""Test BigQueryToolConfig with invalid application name."""
with pytest.raises(
ValueError,
match="Application name should not contain spaces.",
):
BigQueryToolConfig(application_name="my agent")