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
37 changes: 8 additions & 29 deletions src/sagemaker/hyperpod/cli/commands/ray_dashboard_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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}",
Expand All @@ -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(
Expand Down
35 changes: 14 additions & 21 deletions test/unit_tests/cli/test_ray_dashboard_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {
Expand All @@ -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',
Expand All @@ -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',
Expand All @@ -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',
Expand All @@ -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',
Expand All @@ -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(
Expand All @@ -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',
Expand All @@ -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(
Expand All @@ -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',
Expand All @@ -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(
Expand All @@ -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',
Expand Down
Loading