Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 23 additions & 9 deletions google/cloud/dataproc_spark_connect/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -593,20 +593,33 @@ def _get_dataproc_config(self):
self._check_python_version_compatibility(
dataproc_config.runtime_config.version
)

# Use local variable to improve readability of deeply nested attribute access
exec_config = dataproc_config.environment_config.execution_config

# Set service account from environment if not already set
if (
not dataproc_config.environment_config.execution_config.authentication_config.user_workload_authentication_type
and "DATAPROC_SPARK_CONNECT_AUTH_TYPE" in os.environ
):
dataproc_config.environment_config.execution_config.authentication_config.user_workload_authentication_type = AuthenticationConfig.AuthenticationType[
os.getenv("DATAPROC_SPARK_CONNECT_AUTH_TYPE")
]
if (
not dataproc_config.environment_config.execution_config.service_account
not exec_config.service_account
and "DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT" in os.environ
):
dataproc_config.environment_config.execution_config.service_account = os.getenv(
exec_config.service_account = os.getenv(
"DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT"
)

# Auto-set authentication type to SERVICE_ACCOUNT when service account is provided
if exec_config.service_account:
# When service account is provided, explicitly set auth type to SERVICE_ACCOUNT
exec_config.authentication_config.user_workload_authentication_type = (
AuthenticationConfig.AuthenticationType.SERVICE_ACCOUNT
)
elif (
not exec_config.authentication_config.user_workload_authentication_type
and "DATAPROC_SPARK_CONNECT_AUTH_TYPE" in os.environ
):
# Only set auth type from environment if no service account is present
exec_config.authentication_config.user_workload_authentication_type = AuthenticationConfig.AuthenticationType[
os.getenv("DATAPROC_SPARK_CONNECT_AUTH_TYPE")
]
if (
not dataproc_config.environment_config.execution_config.subnetwork_uri
and "DATAPROC_SPARK_CONNECT_SUBNET" in os.environ
Expand Down Expand Up @@ -673,6 +686,7 @@ def _get_dataproc_config(self):
f"DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE is set to an invalid value:"
f" {default_datasource}. Supported value is 'bigquery'."
)

return dataproc_config

def _check_python_version_compatibility(self, runtime_version):
Expand Down
59 changes: 59 additions & 0 deletions tests/unit/test_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -1863,6 +1863,65 @@ def test_builder_pattern_environment_config(
)
self.stopSession(mock_session_controller_client_instance, session)

@mock.patch("google.auth.default")
@mock.patch("google.cloud.dataproc_v1.SessionControllerClient")
@mock.patch("pyspark.sql.connect.client.SparkConnectClient.config")
@mock.patch(
"google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id"
)
@mock.patch(
"google.cloud.dataproc_spark_connect.session.is_s8s_session_active"
)
def test_service_account_sets_auth_type_automatically(
self,
mock_is_s8s_session_active,
mock_dataproc_session_id,
mock_client_config,
mock_session_controller_client,
mock_credentials,
):
"""Test that setting a service account automatically sets auth type to SERVICE_ACCOUNT."""
session = None
mock_session_controller_client_instance = (
self._setup_session_creation_mocks(
mock_is_s8s_session_active,
mock_dataproc_session_id,
mock_client_config,
mock_session_controller_client,
mock_credentials,
)
)

try:
session = DataprocSparkSession.builder.serviceAccount(
"test-service@project.iam.gserviceaccount.com"
).getOrCreate()

# Verify the session was created with the correct authentication config
create_session_request = mock_session_controller_client_instance.create_session.call_args[
0
][
0
]
exec_config = (
create_session_request.session.environment_config.execution_config
)
self.assertEqual(
exec_config.service_account,
"test-service@project.iam.gserviceaccount.com",
)
# Verify that authentication type is automatically set to SERVICE_ACCOUNT
self.assertEqual(
exec_config.authentication_config.user_workload_authentication_type,
AuthenticationConfig.AuthenticationType.SERVICE_ACCOUNT,
)

finally:
mock_session_controller_client_instance.terminate_session.return_value = (
mock.Mock()
)
self.stopSession(mock_session_controller_client_instance, session)

@mock.patch("google.auth.default")
@mock.patch("google.cloud.dataproc_v1.SessionControllerClient")
@mock.patch("pyspark.sql.connect.client.SparkConnectClient.config")
Expand Down