diff --git a/superset/db_engine_specs/bigquery.py b/superset/db_engine_specs/bigquery.py index 42561f05cbb..97e331b1be3 100644 --- a/superset/db_engine_specs/bigquery.py +++ b/superset/db_engine_specs/bigquery.py @@ -728,15 +728,17 @@ class BigQueryEngineSpec(BaseEngineSpec): # pylint: disable=too-many-public-met "Could not import libraries needed to connect to BigQuery." ) + project: str | None = engine.url.host or None + if credentials_info := engine.dialect.credentials_info: credentials = service_account.Credentials.from_service_account_info( credentials_info ) - return bigquery.Client(credentials=credentials) + return bigquery.Client(credentials=credentials, project=project) try: credentials = google.auth.default()[0] - return bigquery.Client(credentials=credentials) + return bigquery.Client(credentials=credentials, project=project) except google.auth.exceptions.DefaultCredentialsError as ex: raise SupersetDBAPIConnectionError( "The database credentials could not be found." diff --git a/tests/unit_tests/db_engine_specs/test_bigquery.py b/tests/unit_tests/db_engine_specs/test_bigquery.py index 69a9923e995..a8402bc8bb4 100644 --- a/tests/unit_tests/db_engine_specs/test_bigquery.py +++ b/tests/unit_tests/db_engine_specs/test_bigquery.py @@ -430,6 +430,78 @@ def test_get_default_catalog(mocker: MockerFixture) -> None: assert BigQueryEngineSpec.get_default_catalog(database) == "project" +@pytest.mark.parametrize( + ("sqlalchemy_uri", "schema", "expected_project"), + [ + ("bigquery://uri-project", None, "uri-project"), + ("bigquery:///uri-project", None, "uri-project"), + ("bigquery://", "dataset_name", None), + ], +) +def test_get_client_resolves_uri_project_with_service_account_credentials( + mocker: MockerFixture, + sqlalchemy_uri: str, + schema: str | None, + expected_project: str | None, +) -> None: + """Test that service-account clients use the project from the engine URI.""" + from superset.db_engine_specs.bigquery import BigQueryEngineSpec + + credentials_info = {"project_id": "credential-project"} + credentials = mock.Mock() + engine = mock.MagicMock() + engine.url = BigQueryEngineSpec.adjust_engine_params( + make_url(sqlalchemy_uri), {}, schema=schema + )[0] + engine.dialect.credentials_info = credentials_info + create_credentials = mocker.patch( + "superset.db_engine_specs.bigquery.service_account.Credentials." + "from_service_account_info", + return_value=credentials, + ) + client = mocker.patch("superset.db_engine_specs.bigquery.bigquery.Client") + + BigQueryEngineSpec._get_client(engine, mock.Mock()) + + create_credentials.assert_called_once_with(credentials_info) + client.assert_called_once_with(credentials=credentials, project=expected_project) + + +@pytest.mark.parametrize( + ("sqlalchemy_uri", "schema", "expected_project"), + [ + ("bigquery://uri-project", None, "uri-project"), + ("bigquery:///uri-project", None, "uri-project"), + ("bigquery://", "dataset_name", None), + ], +) +def test_get_client_resolves_uri_project_with_application_default_credentials( + mocker: MockerFixture, + sqlalchemy_uri: str, + schema: str | None, + expected_project: str | None, +) -> None: + """Test that ADC clients use the project from the engine URI.""" + from superset.db_engine_specs.bigquery import BigQueryEngineSpec + + credentials = mock.Mock() + engine = mock.MagicMock() + engine.url = BigQueryEngineSpec.adjust_engine_params( + make_url(sqlalchemy_uri), {}, schema=schema + )[0] + engine.dialect.credentials_info = None + get_default_credentials = mocker.patch( + "superset.db_engine_specs.bigquery.google.auth.default", + return_value=(credentials, "credential-project"), + ) + client = mocker.patch("superset.db_engine_specs.bigquery.bigquery.Client") + + BigQueryEngineSpec._get_client(engine, mock.Mock()) + + get_default_credentials.assert_called_once_with() + client.assert_called_once_with(credentials=credentials, project=expected_project) + + def test_get_time_partition_column_uses_catalog_in_table_reference( mocker: MockerFixture, ) -> None: