Skip to content
Merged
Show file tree
Hide file tree
Changes from 8 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
29 changes: 22 additions & 7 deletions google/cloud/dataproc_spark_connect/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -593,20 +593,34 @@ def _get_dataproc_config(self):
self._check_python_version_compatibility(
dataproc_config.runtime_config.version
)
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")
]

# Set service account from environment if not already set
if (
not dataproc_config.environment_config.execution_config.service_account
and "DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT" in os.environ
):
dataproc_config.environment_config.execution_config.service_account = os.getenv(
"DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT"
)

# Auto-set authentication type to SERVICE_ACCOUNT when service account is provided
service_account = (
dataproc_config.environment_config.execution_config.service_account
)

if service_account:
# When service account is provided, explicitly set auth type to SERVICE_ACCOUNT
dataproc_config.environment_config.execution_config.authentication_config.user_workload_authentication_type = (
AuthenticationConfig.AuthenticationType.SERVICE_ACCOUNT
)
elif (
not dataproc_config.environment_config.execution_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
dataproc_config.environment_config.execution_config.authentication_config.user_workload_authentication_type = AuthenticationConfig.AuthenticationType[
os.getenv("DATAPROC_SPARK_CONNECT_AUTH_TYPE")
]
Comment thread
fangyh20 marked this conversation as resolved.
if (
not dataproc_config.environment_config.execution_config.subnetwork_uri
and "DATAPROC_SPARK_CONNECT_SUBNET" in os.environ
Expand Down Expand Up @@ -673,6 +687,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
56 changes: 56 additions & 0 deletions tests/unit/test_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -1863,6 +1863,62 @@ 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
]
self.assertEqual(
create_session_request.session.environment_config.execution_config.service_account,
"test-service@project.iam.gserviceaccount.com",
)
# Verify that authentication type is automatically set to SERVICE_ACCOUNT
self.assertEqual(
create_session_request.session.environment_config.execution_config.authentication_config.user_workload_authentication_type,
AuthenticationConfig.AuthenticationType.SERVICE_ACCOUNT,
)
Comment thread
fangyh20 marked this conversation as resolved.

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