From 289941d7ff42bdd5862213848e66b533b2bda10b Mon Sep 17 00:00:00 2001 From: Zhiwei Lin Date: Fri, 26 Sep 2025 13:31:40 -0700 Subject: [PATCH 1/2] fix pysparkvalue error --- google/cloud/dataproc_spark_connect/session.py | 2 +- tests/unit/test_session.py | 7 +++++++ 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/google/cloud/dataproc_spark_connect/session.py b/google/cloud/dataproc_spark_connect/session.py index 09b9f624..90632248 100644 --- a/google/cloud/dataproc_spark_connect/session.py +++ b/google/cloud/dataproc_spark_connect/session.py @@ -519,7 +519,7 @@ def _get_exiting_active_session( self, ) -> Optional["DataprocSparkSession"]: s8s_session_id = DataprocSparkSession._active_s8s_session_id - session_name = f"projects/{self._project_id}/locations/{self._region}/sessions/{s8s_session_id}" + session_name = f"sc://projects/{self._project_id}/locations/{self._region}/sessions/{s8s_session_id}" session_response = None session = None if s8s_session_id is not None: diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 6c10e0ed..fd7cffa9 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -2313,6 +2313,13 @@ def test_create_session_with_client_environment_label( env_label ) + # If the environment has a subnet configured, add it to the expected request. + subnet = os.getenv("DATAPROC_SPARK_CONNECT_SUBNET") + if subnet: + expected_request.session.environment_config.execution_config.subnetwork_uri = ( + subnet + ) + try: # Reset singleton state before each subtest run DataprocSparkSession._active_s8s_session_id = None From ed53c587a7293157caf3e50e7c736259486b50c2 Mon Sep 17 00:00:00 2001 From: Zhiwei Lin Date: Mon, 29 Sep 2025 15:49:53 -0700 Subject: [PATCH 2/2] fix the stop session error --- .../cloud/dataproc_spark_connect/session.py | 18 ++++++++ tests/unit/test_session.py | 41 +++++++++++++++++++ 2 files changed, 59 insertions(+) diff --git a/google/cloud/dataproc_spark_connect/session.py b/google/cloud/dataproc_spark_connect/session.py index 90632248..297c40a3 100644 --- a/google/cloud/dataproc_spark_connect/session.py +++ b/google/cloud/dataproc_spark_connect/session.py @@ -444,6 +444,24 @@ def create_session_pbar(): logger.error( f"Exception while writing active session to file {file_path}, {e}" ) + except KeyboardInterrupt: + # Handle user interruption during session creation + stop_create_session_pbar_event.set() + if create_session_pbar_thread.is_alive(): + create_session_pbar_thread.join() + + logger.warning( + "Session creation interrupted by user. Terminating session..." + ) + terminate_s8s_session( + self._project_id, + self._region, + session_id, + self._client_options, + ) + DataprocSparkSession._active_s8s_session_id = None + DataprocSparkSession._active_session_uses_custom_id = False + raise except (InvalidArgument, PermissionDenied) as e: stop_create_session_pbar_event.set() if create_session_pbar_thread.is_alive(): diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index fd7cffa9..a3558767 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -1362,6 +1362,47 @@ def test_create_session_with_invalid_notebook_id( ) self.stopSession(mock_session_controller_client_instance, session) + @mock.patch("google.auth.default") + @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") + @mock.patch( + "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + ) + @mock.patch( + "google.cloud.dataproc_spark_connect.session.terminate_s8s_session" + ) # Mock terminate + def test_create_session_keyboard_interrupt( + self, + mock_terminate, + mock_dataproc_session_id, + mock_session_controller_client, + mock_credentials, + ): + """Test that KeyboardInterrupt during session creation terminates the session.""" + mock_dataproc_session_id.return_value = "sc-interrupt-test" + mock_session_controller_client_instance = ( + mock_session_controller_client.return_value + ) + mock_operation = mock.Mock() + # Simulate KeyboardInterrupt during operation.result() + mock_operation.result.side_effect = KeyboardInterrupt + mock_session_controller_client_instance.create_session.return_value = ( + mock_operation + ) + cred = mock.MagicMock() + cred.token = "token" + mock_credentials.return_value = (cred, "") + + with self.assertRaises(KeyboardInterrupt): + DataprocSparkSession.builder.getOrCreate() + + # Verify that terminate_s8s_session was called + mock_terminate.assert_called_once() + # Check the arguments passed to terminate_s8s_session + args, _ = mock_terminate.call_args + self.assertEqual(args[0], "test-project") # project_id + self.assertEqual(args[1], "test-region") # region + self.assertEqual(args[2], "sc-interrupt-test") # session_id + @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch(