fix: Set explicit project in the BigQuery client

This change sets an explicit project id in the BigQuery client from the conversation context. Without this the client was trying to set a project from the environment's application default credentials and running into issues where application default credentials is not available.

PiperOrigin-RevId: 772695883
This commit is contained in:
Google Team Member
2025-06-17 18:04:41 -07:00
committed by Copybara-Service
parent c04adaade1
commit 6d174eba30
7 changed files with 315 additions and 20 deletions
+4 -2
View File
@@ -21,13 +21,15 @@ from google.oauth2.credentials import Credentials
USER_AGENT = "adk-bigquery-tool" USER_AGENT = "adk-bigquery-tool"
def get_bigquery_client(*, credentials: Credentials) -> bigquery.Client: def get_bigquery_client(
*, project: str, credentials: Credentials
) -> bigquery.Client:
"""Get a BigQuery client.""" """Get a BigQuery client."""
client_info = google.api_core.client_info.ClientInfo(user_agent=USER_AGENT) client_info = google.api_core.client_info.ClientInfo(user_agent=USER_AGENT)
bigquery_client = bigquery.Client( bigquery_client = bigquery.Client(
credentials=credentials, client_info=client_info project=project, credentials=credentials, client_info=client_info
) )
return bigquery_client return bigquery_client
+12 -4
View File
@@ -42,7 +42,9 @@ def list_dataset_ids(project_id: str, credentials: Credentials) -> list[str]:
'bbc_news'] 'bbc_news']
""" """
try: try:
bq_client = client.get_bigquery_client(credentials=credentials) bq_client = client.get_bigquery_client(
project=project_id, credentials=credentials
)
datasets = [] datasets = []
for dataset in bq_client.list_datasets(project_id): for dataset in bq_client.list_datasets(project_id):
@@ -106,7 +108,9 @@ def get_dataset_info(
} }
""" """
try: try:
bq_client = client.get_bigquery_client(credentials=credentials) bq_client = client.get_bigquery_client(
project=project_id, credentials=credentials
)
dataset = bq_client.get_dataset( dataset = bq_client.get_dataset(
bigquery.DatasetReference(project_id, dataset_id) bigquery.DatasetReference(project_id, dataset_id)
) )
@@ -137,7 +141,9 @@ def list_table_ids(
'local_data_for_better_health_county_data'] 'local_data_for_better_health_county_data']
""" """
try: try:
bq_client = client.get_bigquery_client(credentials=credentials) bq_client = client.get_bigquery_client(
project=project_id, credentials=credentials
)
tables = [] tables = []
for table in bq_client.list_tables( for table in bq_client.list_tables(
@@ -251,7 +257,9 @@ def get_table_info(
} }
""" """
try: try:
bq_client = client.get_bigquery_client(credentials=credentials) bq_client = client.get_bigquery_client(
project=project_id, credentials=credentials
)
return bq_client.get_table( return bq_client.get_table(
bigquery.TableReference( bigquery.TableReference(
bigquery.DatasetReference(project_id, dataset_id), table_id bigquery.DatasetReference(project_id, dataset_id), table_id
+3 -1
View File
@@ -72,7 +72,9 @@ def execute_sql(
""" """
try: try:
bq_client = client.get_bigquery_client(credentials=credentials) bq_client = client.get_bigquery_client(
project=project_id, credentials=credentials
)
if not config or config.write_mode == WriteMode.BLOCKED: if not config or config.write_mode == WriteMode.BLOCKED:
query_job = bq_client.query( query_job = bq_client.query(
query, query,
@@ -0,0 +1,125 @@
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from __future__ import annotations
import os
from unittest import mock
from google.adk.tools.bigquery.client import get_bigquery_client
from google.auth.exceptions import DefaultCredentialsError
from google.oauth2.credentials import Credentials
import pytest
def test_bigquery_client_project():
"""Test BigQuery client project."""
# Trigger the BigQuery client creation
client = get_bigquery_client(
project="test-gcp-project",
credentials=mock.create_autospec(Credentials, instance=True),
)
# Verify that the client has the desired project set
assert client.project == "test-gcp-project"
def test_bigquery_client_project_set_explicit():
"""Test BigQuery client creation does not invoke default auth."""
# Let's simulate that no environment variables are set, so that any project
# set in there does not interfere with this test
with mock.patch.dict(os.environ, {}, clear=True):
with mock.patch("google.auth.default", autospec=True) as mock_default_auth:
# Simulate exception from default auth
mock_default_auth.side_effect = DefaultCredentialsError(
"Your default credentials were not found"
)
# Trigger the BigQuery client creation
client = get_bigquery_client(
project="test-gcp-project",
credentials=mock.create_autospec(Credentials, instance=True),
)
# If we are here that already means client creation did not call default
# auth (otherwise we would have run into DefaultCredentialsError set
# above). For the sake of explicitness, trivially assert that the default
# auth was not called, and yet the project was set correctly
mock_default_auth.assert_not_called()
assert client.project == "test-gcp-project"
def test_bigquery_client_project_set_with_default_auth():
"""Test BigQuery client creation invokes default auth to set the project."""
# Let's simulate that no environment variables are set, so that any project
# set in there does not interfere with this test
with mock.patch.dict(os.environ, {}, clear=True):
with mock.patch("google.auth.default", autospec=True) as mock_default_auth:
# Simulate credentials
mock_creds = mock.create_autospec(Credentials, instance=True)
# Simulate output of the default auth
mock_default_auth.return_value = (mock_creds, "test-gcp-project")
# Trigger the BigQuery client creation
client = get_bigquery_client(
project=None,
credentials=mock_creds,
)
# Verify that default auth was called once to set the client project
mock_default_auth.assert_called_once()
assert client.project == "test-gcp-project"
def test_bigquery_client_project_set_with_env():
"""Test BigQuery client creation sets the project from environment variable."""
# Let's simulate the project set in environment variables
with mock.patch.dict(
os.environ, {"GOOGLE_CLOUD_PROJECT": "test-gcp-project"}, clear=True
):
with mock.patch("google.auth.default", autospec=True) as mock_default_auth:
# Simulate exception from default auth
mock_default_auth.side_effect = DefaultCredentialsError(
"Your default credentials were not found"
)
# Trigger the BigQuery client creation
client = get_bigquery_client(
project=None,
credentials=mock.create_autospec(Credentials, instance=True),
)
# If we are here that already means client creation did not call default
# auth (otherwise we would have run into DefaultCredentialsError set
# above). For the sake of explicitness, trivially assert that the default
# auth was not called, and yet the project was set correctly
mock_default_auth.assert_not_called()
assert client.project == "test-gcp-project"
def test_bigquery_client_user_agent():
"""Test BigQuery client 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),
)
# 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 client_info_arg.user_agent == "adk-bigquery-tool"
@@ -0,0 +1,122 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from __future__ import annotations
import os
from unittest import mock
from google.adk.tools.bigquery import metadata_tool
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."""
project = "my_project_id"
mock_credentials = mock.create_autospec(Credentials, instance=True)
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
mock_default_auth.side_effect = DefaultCredentialsError(
"Your default credentials were not found"
)
mock_list_datasets.return_value = [
bigquery.DatasetReference(project, "dataset1"),
bigquery.DatasetReference(project, "dataset2"),
]
result = metadata_tool.list_dataset_ids(project, mock_credentials)
assert result == ["dataset1", "dataset2"]
mock_default_auth.assert_not_called()
@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."""
mock_credentials = mock.create_autospec(Credentials, instance=True)
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
mock_default_auth.side_effect = DefaultCredentialsError(
"Your default credentials were not found"
)
mock_get_dataset.return_value = mock.create_autospec(
Credentials, instance=True
)
result = metadata_tool.get_dataset_info(
"my_project_id", "my_dataset_id", mock_credentials
)
assert result != {
"status": "ERROR",
"error_details": "Your default credentials were not found",
}
mock_default_auth.assert_not_called()
@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."""
project = "my_project_id"
dataset = "my_dataset_id"
dataset_ref = bigquery.DatasetReference(project, dataset)
mock_credentials = mock.create_autospec(Credentials, instance=True)
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
mock_default_auth.side_effect = DefaultCredentialsError(
"Your default credentials were not found"
)
mock_list_tables.return_value = [
bigquery.TableReference(dataset_ref, "table1"),
bigquery.TableReference(dataset_ref, "table2"),
]
result = metadata_tool.list_table_ids(project, dataset, mock_credentials)
assert result == ["table1", "table2"]
mock_default_auth.assert_not_called()
@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."""
mock_credentials = mock.create_autospec(Credentials, instance=True)
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
mock_default_auth.side_effect = DefaultCredentialsError(
"Your default credentials were not found"
)
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
)
assert result != {
"status": "ERROR",
"error_details": "Your default credentials were not found",
}
mock_default_auth.assert_not_called()
@@ -14,6 +14,7 @@
from __future__ import annotations from __future__ import annotations
import os
import textwrap import textwrap
from typing import Optional from typing import Optional
from unittest import mock from unittest import mock
@@ -24,6 +25,7 @@ from google.adk.tools.bigquery import BigQueryToolset
from google.adk.tools.bigquery.config import BigQueryToolConfig from google.adk.tools.bigquery.config import BigQueryToolConfig
from google.adk.tools.bigquery.config import WriteMode from google.adk.tools.bigquery.config import WriteMode
from google.adk.tools.bigquery.query_tool import execute_sql from google.adk.tools.bigquery.query_tool import execute_sql
from google.auth.exceptions import DefaultCredentialsError
from google.cloud import bigquery from google.cloud import bigquery
from google.oauth2.credentials import Credentials from google.oauth2.credentials import Credentials
import pytest import pytest
@@ -227,14 +229,8 @@ async def test_execute_sql_declaration_write(tool_config):
@pytest.mark.parametrize( @pytest.mark.parametrize(
("write_mode",), ("write_mode",),
[ [
pytest.param( pytest.param(WriteMode.BLOCKED, id="blocked"),
WriteMode.BLOCKED, pytest.param(WriteMode.ALLOWED, id="allowed"),
id="blocked",
),
pytest.param(
WriteMode.ALLOWED,
id="allowed",
),
], ],
) )
def test_execute_sql_select_stmt(write_mode): def test_execute_sql_select_stmt(write_mode):
@@ -279,7 +275,7 @@ def test_execute_sql_select_stmt(write_mode):
], ],
) )
def test_execute_sql_non_select_stmt_write_allowed(query, statement_type): def test_execute_sql_non_select_stmt_write_allowed(query, statement_type):
"""Test execute_sql tool for SELECT query when writes are blocked.""" """Test execute_sql tool for non-SELECT query when writes are blocked."""
project = "my_project" project = "my_project"
query_result = [] query_result = []
credentials = mock.create_autospec(Credentials, instance=True) credentials = mock.create_autospec(Credentials, instance=True)
@@ -318,7 +314,7 @@ def test_execute_sql_non_select_stmt_write_allowed(query, statement_type):
], ],
) )
def test_execute_sql_non_select_stmt_write_blocked(query, statement_type): def test_execute_sql_non_select_stmt_write_blocked(query, statement_type):
"""Test execute_sql tool for SELECT query when writes are blocked.""" """Test execute_sql tool for non-SELECT query when writes are blocked."""
project = "my_project" project = "my_project"
query_result = [] query_result = []
credentials = mock.create_autospec(Credentials, instance=True) credentials = mock.create_autospec(Credentials, instance=True)
@@ -342,3 +338,45 @@ def test_execute_sql_non_select_stmt_write_blocked(query, statement_type):
"status": "ERROR", "status": "ERROR",
"error_details": "Read-only mode only supports SELECT statements.", "error_details": "Read-only mode only supports SELECT statements.",
} }
@pytest.mark.parametrize(
("write_mode",),
[
pytest.param(WriteMode.BLOCKED, id="blocked"),
pytest.param(WriteMode.ALLOWED, id="allowed"),
],
)
@mock.patch.dict(os.environ, {}, clear=True)
@mock.patch("google.cloud.bigquery.Client.query_and_wait", autospec=True)
@mock.patch("google.cloud.bigquery.Client.query", autospec=True)
@mock.patch("google.auth.default", autospec=True)
def test_execute_sql_no_default_auth(
mock_default_auth, mock_query, mock_query_and_wait, write_mode
):
"""Test execute_sql tool invocation does not involve calling default auth."""
project = "my_project"
query = "SELECT 123 AS num"
statement_type = "SELECT"
query_result = [{"num": 123}]
credentials = mock.create_autospec(Credentials, instance=True)
tool_config = BigQueryToolConfig(write_mode=write_mode)
# Simulate the behavior of default auth - on purpose throw exception when
# the default auth is called
mock_default_auth.side_effect = DefaultCredentialsError(
"Your default credentials were not found"
)
# Simulate the result of query API
query_job = mock.create_autospec(bigquery.QueryJob)
query_job.statement_type = statement_type
mock_query.return_value = query_job
# Simulate the result of query_and_wait API
mock_query_and_wait.return_value = query_result
# Test the tool worked without invoking default auth
result = execute_sql(project, query, credentials, tool_config)
assert result == {"status": "SUCCESS", "rows": query_result}
mock_default_auth.assert_not_called()
@@ -96,9 +96,7 @@ async def test_bigquery_toolset_tools_selective(selected_tools):
], ],
) )
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_bigquery_toolset_unknown_tool_raises( async def test_bigquery_toolset_unknown_tool(selected_tools, returned_tools):
selected_tools, returned_tools
):
"""Test BigQuery toolset with filter. """Test BigQuery toolset with filter.
This test verifies the behavior of the BigQuery toolset when filter is This test verifies the behavior of the BigQuery toolset when filter is