From 5c2f406907395ad10a6f9e59402e701facbdc973 Mon Sep 17 00:00:00 2001 From: Joyjit Chatterjee Date: Tue, 1 Sep 2026 23:07:52 +0000 Subject: [PATCH] fix: Load K8s credentials for create ray-dashboard-connection through load_kubeconfig --- .../cli/commands/ray_dashboard_connection.py | 37 ++++--------------- .../cli/test_ray_dashboard_connection.py | 35 +++++++----------- 2 files changed, 22 insertions(+), 50 deletions(-) diff --git a/src/sagemaker/hyperpod/cli/commands/ray_dashboard_connection.py b/src/sagemaker/hyperpod/cli/commands/ray_dashboard_connection.py index 0eedb2f2..6795dc8c 100644 --- a/src/sagemaker/hyperpod/cli/commands/ray_dashboard_connection.py +++ b/src/sagemaker/hyperpod/cli/commands/ray_dashboard_connection.py @@ -27,36 +27,15 @@ from sagemaker.hyperpod.common.cli_decorators import handle_cli_exceptions -def _get_eks_api_client(): - """Load kubeconfig and create an authenticated API client. +def _load_kube_config(): + """Load kubeconfig so the default ApiClient is authenticated. - Works around a kubernetes-client issue where exec-based tokens - are not properly forwarded in the Authorization header. + Uses the library's default client (as HPSpace does) rather than copying + the token out of Configuration.api_key. The api_key entry is named + differently across kubernetes-client releases ("authorization" in 36.0.0, + "BearerToken" in 36.0.3+), so reading it by name is version-fragile. """ config.load_kube_config() - configuration = client.Configuration.get_default_copy() - - # Extract the token from the exec provider - token = None - if configuration.api_key and "authorization" in configuration.api_key: - token_value = configuration.api_key["authorization"] - prefix = "Bearer " - if token_value.startswith(prefix): - token = token_value.removeprefix(prefix) - else: - token = token_value - - # Clear api_key to avoid double-auth conflicts - configuration.api_key = {} - configuration.api_key_prefix = {} - - if token: - return client.ApiClient( - configuration, - header_name="Authorization", - header_value=f"Bearer {token}", - ) - return client.ApiClient(configuration) @click.command("ray-dashboard-connection") @@ -66,7 +45,7 @@ def _get_eks_api_client(): @handle_cli_exceptions() def create_ray_dashboard_connection(cluster_name, namespace): """Create a RayDashboardConnection to get a dashboard URL for a RayCluster.""" - api_client = _get_eks_api_client() + _load_kube_config() body = { "apiVersion": f"{RAY_DASHBOARD_CONNECTION_GROUP}/{RAY_DASHBOARD_CONNECTION_VERSION}", @@ -79,7 +58,7 @@ def create_ray_dashboard_connection(cluster_name, namespace): }, } - api = client.CustomObjectsApi(api_client) + api = client.CustomObjectsApi() try: result = api.create_namespaced_custom_object( diff --git a/test/unit_tests/cli/test_ray_dashboard_connection.py b/test/unit_tests/cli/test_ray_dashboard_connection.py index 2f595d89..4b8fc8f7 100644 --- a/test/unit_tests/cli/test_ray_dashboard_connection.py +++ b/test/unit_tests/cli/test_ray_dashboard_connection.py @@ -13,9 +13,9 @@ class TestRayDashboardConnectionCommand: def setup_method(self): self.runner = CliRunner() - @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._get_eks_api_client') + @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._load_kube_config') @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection.client.CustomObjectsApi') - def test_create_success_returns_url(self, mock_custom_objects_api_class, mock_get_client): + def test_create_success_returns_url(self, mock_custom_objects_api_class, mock_load_config): """Test successful creation returns the connection URL""" mock_api = Mock() mock_api.create_namespaced_custom_object.return_value = { @@ -24,7 +24,6 @@ def test_create_success_returns_url(self, mock_custom_objects_api_class, mock_ge } } mock_custom_objects_api_class.return_value = mock_api - mock_get_client.return_value = Mock() result = self.runner.invoke(create_ray_dashboard_connection, [ '--cluster-name', 'my-raycluster', @@ -46,16 +45,15 @@ def test_create_success_returns_url(self, mock_custom_objects_api_class, mock_ge }, ) - @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._get_eks_api_client') + @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._load_kube_config') @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection.client.CustomObjectsApi') - def test_create_default_namespace(self, mock_custom_objects_api_class, mock_get_client): + def test_create_default_namespace(self, mock_custom_objects_api_class, mock_load_config): """Test namespace defaults to 'default' when not specified""" mock_api = Mock() mock_api.create_namespaced_custom_object.return_value = { "status": {"connectionUrl": "https://example.com/dashboard"} } mock_custom_objects_api_class.return_value = mock_api - mock_get_client.return_value = Mock() result = self.runner.invoke(create_ray_dashboard_connection, [ '--cluster-name', 'my-raycluster', @@ -66,16 +64,15 @@ def test_create_default_namespace(self, mock_custom_objects_api_class, mock_get_ call_kwargs = mock_api.create_namespaced_custom_object.call_args[1] assert call_kwargs["namespace"] == "default" - @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._get_eks_api_client') + @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._load_kube_config') @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection.client.CustomObjectsApi') - def test_create_empty_url_raises_error(self, mock_custom_objects_api_class, mock_get_client): + def test_create_empty_url_raises_error(self, mock_custom_objects_api_class, mock_load_config): """Test that empty connectionUrl raises an error""" mock_api = Mock() mock_api.create_namespaced_custom_object.return_value = { "status": {"connectionUrl": ""} } mock_custom_objects_api_class.return_value = mock_api - mock_get_client.return_value = Mock() result = self.runner.invoke(create_ray_dashboard_connection, [ '--cluster-name', 'my-raycluster', @@ -86,16 +83,15 @@ def test_create_empty_url_raises_error(self, mock_custom_objects_api_class, mock assert "Failed to get dashboard URL" in result.output assert "contact your cluster administrator" in result.output - @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._get_eks_api_client') + @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._load_kube_config') @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection.client.CustomObjectsApi') - def test_create_no_status_raises_error(self, mock_custom_objects_api_class, mock_get_client): + def test_create_no_status_raises_error(self, mock_custom_objects_api_class, mock_load_config): """Test that missing status raises an error""" mock_api = Mock() mock_api.create_namespaced_custom_object.return_value = { "metadata": {"name": "generated-name"} } mock_custom_objects_api_class.return_value = mock_api - mock_get_client.return_value = Mock() result = self.runner.invoke(create_ray_dashboard_connection, [ '--cluster-name', 'my-raycluster', @@ -104,9 +100,9 @@ def test_create_no_status_raises_error(self, mock_custom_objects_api_class, mock assert result.exit_code != 0 assert "Failed to get dashboard URL" in result.output - @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._get_eks_api_client') + @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._load_kube_config') @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection.client.CustomObjectsApi') - def test_create_404_api_not_installed(self, mock_custom_objects_api_class, mock_get_client): + def test_create_404_api_not_installed(self, mock_custom_objects_api_class, mock_load_config): """Test 404 when operator is not installed shows install instructions""" mock_api = Mock() mock_api.create_namespaced_custom_object.side_effect = ApiException( @@ -123,7 +119,6 @@ def test_create_404_api_not_installed(self, mock_custom_objects_api_class, mock_ '"details":{"group":"connection.access.sagemaker.amazonaws.com","kind":"raydashboardconnections"}}' ) mock_custom_objects_api_class.return_value = mock_api - mock_get_client.return_value = Mock() result = self.runner.invoke(create_ray_dashboard_connection, [ '--cluster-name', 'my-raycluster', @@ -135,9 +130,9 @@ def test_create_404_api_not_installed(self, mock_custom_objects_api_class, mock_ assert "hyperpod-ray-endpoint-operator" in result.output @patch('sagemaker.hyperpod.common.cli_decorators._namespace_exists', return_value=True) - @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._get_eks_api_client') + @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._load_kube_config') @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection.client.CustomObjectsApi') - def test_create_404_namespace_not_found(self, mock_custom_objects_api_class, mock_get_client, mock_ns_exists): + def test_create_404_namespace_not_found(self, mock_custom_objects_api_class, mock_load_config, mock_ns_exists): """Test 404 for missing namespace shows raw error""" mock_api = Mock() mock_api.create_namespaced_custom_object.side_effect = ApiException( @@ -147,7 +142,6 @@ def test_create_404_namespace_not_found(self, mock_custom_objects_api_class, moc ) mock_api.create_namespaced_custom_object.side_effect.body = '{"message":"namespaces not-exists not found"}' mock_custom_objects_api_class.return_value = mock_api - mock_get_client.return_value = Mock() result = self.runner.invoke(create_ray_dashboard_connection, [ '--cluster-name', 'my-raycluster', @@ -158,9 +152,9 @@ def test_create_404_namespace_not_found(self, mock_custom_objects_api_class, moc assert "not found" in result.output.lower() or "not-exists" in result.output @patch('sagemaker.hyperpod.common.cli_decorators._namespace_exists', return_value=True) - @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._get_eks_api_client') + @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection._load_kube_config') @patch('sagemaker.hyperpod.cli.commands.ray_dashboard_connection.client.CustomObjectsApi') - def test_create_403_raises_exception(self, mock_custom_objects_api_class, mock_get_client, mock_ns_exists): + def test_create_403_raises_exception(self, mock_custom_objects_api_class, mock_load_config, mock_ns_exists): """Test 403 forbidden is propagated as an error""" mock_api = Mock() exc = ApiException( @@ -171,7 +165,6 @@ def test_create_403_raises_exception(self, mock_custom_objects_api_class, mock_g exc.body = '{"message":"forbidden"}' mock_api.create_namespaced_custom_object.side_effect = exc mock_custom_objects_api_class.return_value = mock_api - mock_get_client.return_value = Mock() result = self.runner.invoke(create_ray_dashboard_connection, [ '--cluster-name', 'my-raycluster',