mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
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:
committed by
Copybara-Service
parent
f4a8df0ba2
commit
11a2ffe35a
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user