diff --git a/google/cloud/managed_spark_connect/session.py b/google/cloud/managed_spark_connect/session.py
index 76d2700..4f76a49 100644
--- a/google/cloud/managed_spark_connect/session.py
+++ b/google/cloud/managed_spark_connect/session.py
@@ -27,7 +27,7 @@
import tqdm
from packaging import version
from types import MethodType
-from typing import Any, cast, ClassVar, Dict, Iterable, Optional, Union
+from typing import Any, cast, ClassVar, Dict, Iterable, List, Optional, Tuple, Union
from google.api_core import retry
from google.api_core.client_options import ClientOptions
@@ -71,6 +71,21 @@
"https://console.cloud.google.com/dataproc/interactive"
)
+_VSCODE_SESSION_URI_BASE = (
+ "vscode://googlecloudtools.datacloud/dataproc/sessions"
+)
+
+
+def _build_session_details_links(
+ region: Optional[str], project_id: Optional[str], session_id: str
+) -> List[Tuple[str, str]]:
+ console_url = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{region}/{session_id}?project={project_id}"
+ links = [("Managed Spark Session (Cloud Console)", console_url)]
+ if environment.is_vscode():
+ vscode_url = f"{_VSCODE_SESSION_URI_BASE}/{session_id}?project={project_id}&location={region}"
+ links.append(("Managed Spark Session (Data Agent Kit)", vscode_url))
+ return links
+
def _is_valid_label_value(value: str) -> bool:
"""
@@ -506,9 +521,11 @@ def _wait_for_session_available(
)
def _display_session_link_on_creation(self, session_id):
- session_url = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/{session_id}?project={self._project_id}"
- plain_message = (
- f"Creating Managed Spark Connect Session: {session_url}"
+ links = _build_session_details_links(
+ self._region, self._project_id, session_id
+ )
+ plain_message = "Creating Managed Spark Connect Session:\n" + (
+ "\n".join(f" {label}: {url}" for label, url in links)
)
if environment.is_colab_enterprise():
html_element = f"""
@@ -517,10 +534,14 @@ def _display_session_link_on_creation(self, session_id):
"""
else:
+ links_html = "\n".join(
+ f'
{label}
'
+ for label, url in links
+ )
html_element = f"""
"""
self._output_element_or_message(plain_message, html_element)
@@ -573,8 +594,15 @@ def _get_exiting_active_session(
session = ManagedSparkSession._default_session
if session_response is not None:
+ links = _build_session_details_links(
+ self._region, self._project_id, s8s_session_id
+ )
+ links_message = "\n".join(
+ f" {label}: {url}" for label, url in links
+ )
print(
- f"Using existing Managed Spark Session (configuration changes may not be applied): {_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/{s8s_session_id}?project={self._project_id}"
+ "Using existing Managed Spark Session (configuration "
+ f"changes may not be applied):\n{links_message}"
)
self._display_view_session_details_button(s8s_session_id)
if session is None:
@@ -1107,14 +1135,24 @@ def _repr_html_(self) -> str:
No Active Managed Spark Session
"""
- s8s_session = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/{self._active_s8s_session_id}"
- ui = f"{s8s_session}/sparkApplications/applications"
+ session_links = _build_session_details_links(
+ self._region, self._project_id, self._active_s8s_session_id
+ )
+ session_links_html = "\n".join(
+ f'{label}
'
+ for label, url in session_links
+ )
+ ssui_url = (
+ f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/"
+ f"{self._active_s8s_session_id}/sparkApplications/applications"
+ f"?project={self._project_id}"
+ )
return f"""
"""
diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py
index a88bd7a..dbd4811 100644
--- a/tests/unit/test_session.py
+++ b/tests/unit/test_session.py
@@ -1259,6 +1259,136 @@ def test_display_session_link_on_creation_not_colab_enterprise(
self.assertIn("Creating Managed Spark Connect Session", html_output)
self.assertIn("Managed Spark Session", html_output)
+ @mock.patch(
+ "IPython.core.interactiveshell.InteractiveShell.initialized",
+ return_value=True,
+ )
+ @mock.patch("IPython.display.display")
+ def test_display_session_link_on_creation_vscode(
+ self,
+ mock_display,
+ _mock_ipy,
+ ):
+ mock.patch.dict(
+ os.environ,
+ {
+ "VSCODE_PID": "12345",
+ },
+ ).start()
+ ManagedSparkSession.builder._display_session_link_on_creation(
+ "test_session"
+ )
+
+ mock_display.assert_called_once()
+ args, _ = mock_display.call_args
+ html_output = args[0].data
+ self.assertIn("Creating Managed Spark Connect Session", html_output)
+ self.assertIn(
+ f''
+ "Managed Spark Session (Cloud Console)",
+ html_output,
+ )
+ self.assertIn(
+ "vscode://googlecloudtools.datacloud/dataproc/sessions/"
+ "test_session?project=test-project&location=test-region",
+ html_output,
+ )
+ self.assertIn("Managed Spark Session (Data Agent Kit)", html_output)
+
+ @mock.patch.object(ManagedSparkSession, "getActiveSession")
+ @mock.patch(
+ "google.cloud.managed_spark_connect.session.get_active_s8s_session_response"
+ )
+ def test_get_exiting_active_session_prints_vscode_url(
+ self,
+ mock_get_response,
+ mock_get_active_session,
+ ):
+ mock.patch.dict(
+ os.environ,
+ {
+ "VSCODE_PID": "12345",
+ },
+ ).start()
+ mock_get_response.return_value = mock.Mock()
+ mock_get_active_session.return_value = mock.Mock()
+ ManagedSparkSession._active_s8s_session_id = "test_session"
+ self.addCleanup(
+ setattr, ManagedSparkSession, "_active_s8s_session_id", None
+ )
+
+ with mock.patch("builtins.print") as mock_print:
+ ManagedSparkSession.builder._get_exiting_active_session()
+
+ printed = "\n".join(
+ str(call.args[0]) for call in mock_print.call_args_list
+ )
+ self.assertIn("Managed Spark Session (Cloud Console)", printed)
+ self.assertIn(
+ f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/test-region/"
+ "test_session?project=test-project",
+ printed,
+ )
+ self.assertIn("Managed Spark Session (Data Agent Kit)", printed)
+ self.assertIn(
+ "vscode://googlecloudtools.datacloud/dataproc/sessions/"
+ "test_session?project=test-project&location=test-region",
+ printed,
+ )
+
+ def test_repr_html_uses_vscode_url_for_session_link(self):
+ mock.patch.dict(
+ os.environ,
+ {
+ "VSCODE_PID": "12345",
+ },
+ ).start()
+ ManagedSparkSession._project_id = "test-project"
+ ManagedSparkSession._region = "test-region"
+ ManagedSparkSession._active_s8s_session_id = "test_session"
+ self.addCleanup(setattr, ManagedSparkSession, "_project_id", None)
+ self.addCleanup(setattr, ManagedSparkSession, "_region", None)
+ self.addCleanup(
+ setattr, ManagedSparkSession, "_active_s8s_session_id", None
+ )
+
+ html = object.__new__(ManagedSparkSession)._repr_html_()
+
+ self.assertIn(
+ f''
+ "Managed Spark Session (Cloud Console)",
+ html,
+ )
+ self.assertIn(
+ ''
+ "Managed Spark Session (Data Agent Kit)",
+ html,
+ )
+ self.assertIn(
+ f''
+ "Spark UI",
+ html,
+ )
+
+ def test_repr_html_no_vscode_link_when_not_in_vscode(self):
+ os.environ.pop("VSCODE_PID", None)
+ ManagedSparkSession._project_id = "test-project"
+ ManagedSparkSession._region = "test-region"
+ ManagedSparkSession._active_s8s_session_id = "test_session"
+ self.addCleanup(setattr, ManagedSparkSession, "_project_id", None)
+ self.addCleanup(setattr, ManagedSparkSession, "_region", None)
+ self.addCleanup(
+ setattr, ManagedSparkSession, "_active_s8s_session_id", None
+ )
+
+ html = object.__new__(ManagedSparkSession)._repr_html_()
+
+ self.assertNotIn("vscode://", html)
+
def test_is_valid_label_value(self):
# Valid label values
self.assertTrue(_is_valid_label_value("valid-label-123"))