mirror of
https://github.com/NVIDIA/OpenShell.git
synced 2026-10-02 07:34:45 +08:00
* refactor(auth): separate sandbox identity from TLS Signed-off-by: Drew Newberry <anewberry@nvidia.com> * docs(auth): clarify gateway mTLS behavior Signed-off-by: Drew Newberry <anewberry@nvidia.com> * test(auth): include workspace scope in TLS authorization checks Signed-off-by: Drew Newberry <anewberry@nvidia.com> * test(e2e): bound service auth sandbox names for large PIDs Signed-off-by: Drew Newberry <anewberry@nvidia.com> --------- Signed-off-by: Drew Newberry <anewberry@nvidia.com>
271 lines
9.6 KiB
Python
271 lines
9.6 KiB
Python
# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
|
|
"""E2e tests for gateway TLS and mTLS user authentication.
|
|
|
|
TLS accepts CA-only clients so supervisors can use sandbox bearer tokens.
|
|
Health is public; user RPCs require an authenticated user. Presented client
|
|
certificates must be signed by the gateway CA.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import os
|
|
import pathlib
|
|
import subprocess
|
|
import tempfile
|
|
from urllib.parse import urlparse
|
|
|
|
import grpc
|
|
import pytest
|
|
|
|
from openshell._proto import datamodel_pb2, openshell_pb2, openshell_pb2_grpc
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _xdg_config_home() -> pathlib.Path:
|
|
configured = os.environ.get("XDG_CONFIG_HOME")
|
|
if configured:
|
|
return pathlib.Path(configured)
|
|
return pathlib.Path.home() / ".config"
|
|
|
|
|
|
def _resolve_cluster_name() -> str:
|
|
if os.environ.get("OPENSHELL_GATEWAY_ENDPOINT"):
|
|
return os.environ.get("OPENSHELL_GATEWAY", "openshell-e2e-endpoint")
|
|
env_cluster = os.environ.get("OPENSHELL_GATEWAY")
|
|
if env_cluster:
|
|
return env_cluster
|
|
active_file = _xdg_config_home() / "openshell" / "active_gateway"
|
|
return active_file.read_text().strip()
|
|
|
|
|
|
def _cluster_metadata(cluster_name: str) -> dict:
|
|
endpoint = os.environ.get("OPENSHELL_GATEWAY_ENDPOINT")
|
|
if endpoint:
|
|
return {
|
|
"name": cluster_name,
|
|
"gateway_endpoint": endpoint,
|
|
"auth_mode": "plaintext",
|
|
}
|
|
metadata_path = (
|
|
_xdg_config_home() / "openshell" / "gateways" / cluster_name / "metadata.json"
|
|
)
|
|
return json.loads(metadata_path.read_text())
|
|
|
|
|
|
def _mtls_dir(cluster_name: str) -> pathlib.Path:
|
|
return _xdg_config_home() / "openshell" / "gateways" / cluster_name / "mtls"
|
|
|
|
|
|
def _generate_self_signed_cert(
|
|
tmpdir: pathlib.Path,
|
|
) -> tuple[pathlib.Path, pathlib.Path]:
|
|
"""Generate a self-signed cert+key pair that is NOT signed by the cluster CA."""
|
|
cert_path = tmpdir / "rogue.crt"
|
|
key_path = tmpdir / "rogue.key"
|
|
subprocess.run(
|
|
[
|
|
"openssl",
|
|
"req",
|
|
"-x509",
|
|
"-sha256",
|
|
"-nodes",
|
|
"-days",
|
|
"1",
|
|
"-newkey",
|
|
"rsa:2048",
|
|
"-subj",
|
|
"/O=rogue/CN=rogue-client",
|
|
"-keyout",
|
|
str(key_path),
|
|
"-out",
|
|
str(cert_path),
|
|
],
|
|
check=True,
|
|
capture_output=True,
|
|
)
|
|
return cert_path, key_path
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Fixtures
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@pytest.fixture(scope="session")
|
|
def cluster_name() -> str:
|
|
name = _resolve_cluster_name()
|
|
if not name:
|
|
pytest.skip("no active cluster configured")
|
|
return name
|
|
|
|
|
|
@pytest.fixture(scope="session")
|
|
def server_endpoint(cluster_name: str) -> tuple[str, int, str]:
|
|
"""Return (host, port, scheme) for the OpenShell server."""
|
|
metadata = _cluster_metadata(cluster_name)
|
|
parsed = urlparse(metadata["gateway_endpoint"])
|
|
host = parsed.hostname or "127.0.0.1"
|
|
port = parsed.port or (443 if parsed.scheme == "https" else 8080)
|
|
return host, port, parsed.scheme
|
|
|
|
|
|
@pytest.fixture(scope="session")
|
|
def mtls_certs(
|
|
cluster_name: str, server_endpoint: tuple[str, int, str]
|
|
) -> tuple[bytes, bytes, bytes]:
|
|
"""Return (ca_pem, cert_pem, key_pem) for the provisioned mTLS client."""
|
|
_, _, scheme = server_endpoint
|
|
if scheme != "https":
|
|
pytest.skip("server is not using TLS; mTLS tests require an HTTPS endpoint")
|
|
mtls = _mtls_dir(cluster_name)
|
|
ca = (mtls / "ca.crt").read_bytes()
|
|
cert = (mtls / "tls.crt").read_bytes()
|
|
key = (mtls / "tls.key").read_bytes()
|
|
return ca, cert, key
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Tests
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestServerMtlsEnforcement:
|
|
"""Verify TLS trust and the mTLS user authorization boundary."""
|
|
|
|
def test_authenticated_client_succeeds(
|
|
self,
|
|
server_endpoint: tuple[str, int, str],
|
|
mtls_certs: tuple[bytes, bytes, bytes],
|
|
) -> None:
|
|
"""A verified mTLS user can call Health and a protected user RPC."""
|
|
host, port, _ = server_endpoint
|
|
ca, cert, key = mtls_certs
|
|
|
|
credentials = grpc.ssl_channel_credentials(
|
|
root_certificates=ca,
|
|
private_key=key,
|
|
certificate_chain=cert,
|
|
)
|
|
channel = grpc.secure_channel(f"{host}:{port}", credentials)
|
|
try:
|
|
stub = openshell_pb2_grpc.OpenShellStub(channel)
|
|
response = stub.Health(openshell_pb2.HealthRequest(), timeout=10)
|
|
assert response.status == openshell_pb2.SERVICE_STATUS_HEALTHY
|
|
stub.ListSandboxes(
|
|
openshell_pb2.ListSandboxesRequest(
|
|
workspace_scope=datamodel_pb2.WorkspaceSelector(workspace="default")
|
|
),
|
|
timeout=10,
|
|
)
|
|
finally:
|
|
channel.close()
|
|
|
|
def test_ca_only_client_health_succeeds_but_user_rpc_rejected(
|
|
self,
|
|
server_endpoint: tuple[str, int, str],
|
|
mtls_certs: tuple[bytes, bytes, bytes],
|
|
) -> None:
|
|
"""CA-only TLS reaches Health but cannot acquire mTLS user identity."""
|
|
host, port, _ = server_endpoint
|
|
ca, _, _ = mtls_certs
|
|
|
|
# Only provide the CA for server verification -- no client cert/key.
|
|
credentials = grpc.ssl_channel_credentials(root_certificates=ca)
|
|
channel = grpc.secure_channel(f"{host}:{port}", credentials)
|
|
try:
|
|
stub = openshell_pb2_grpc.OpenShellStub(channel)
|
|
response = stub.Health(openshell_pb2.HealthRequest(), timeout=10)
|
|
assert response.status == openshell_pb2.SERVICE_STATUS_HEALTHY
|
|
with pytest.raises(grpc.RpcError) as exc_info:
|
|
stub.ListSandboxes(
|
|
openshell_pb2.ListSandboxesRequest(
|
|
workspace_scope=datamodel_pb2.WorkspaceSelector(
|
|
workspace="default"
|
|
)
|
|
),
|
|
timeout=10,
|
|
)
|
|
assert exc_info.value.code() == grpc.StatusCode.UNAUTHENTICATED
|
|
|
|
# An unverified bearer token must not promote the TLS connection
|
|
# to a user identity either.
|
|
with pytest.raises(grpc.RpcError) as exc_info:
|
|
stub.ListSandboxes(
|
|
openshell_pb2.ListSandboxesRequest(
|
|
workspace_scope=datamodel_pb2.WorkspaceSelector(
|
|
workspace="default"
|
|
)
|
|
),
|
|
metadata=(("authorization", "Bearer invalid-token"),),
|
|
timeout=10,
|
|
)
|
|
assert exc_info.value.code() == grpc.StatusCode.UNAUTHENTICATED
|
|
finally:
|
|
channel.close()
|
|
|
|
def test_wrong_client_cert_rejected(
|
|
self,
|
|
server_endpoint: tuple[str, int, str],
|
|
mtls_certs: tuple[bytes, bytes, bytes],
|
|
) -> None:
|
|
"""A client presenting a cert signed by a different CA is rejected."""
|
|
host, port, _ = server_endpoint
|
|
ca, _, _ = mtls_certs
|
|
|
|
with tempfile.TemporaryDirectory() as tmpdir:
|
|
rogue_cert_path, rogue_key_path = _generate_self_signed_cert(
|
|
pathlib.Path(tmpdir)
|
|
)
|
|
rogue_cert = rogue_cert_path.read_bytes()
|
|
rogue_key = rogue_key_path.read_bytes()
|
|
|
|
credentials = grpc.ssl_channel_credentials(
|
|
root_certificates=ca,
|
|
private_key=rogue_key,
|
|
certificate_chain=rogue_cert,
|
|
)
|
|
channel = grpc.secure_channel(f"{host}:{port}", credentials)
|
|
try:
|
|
stub = openshell_pb2_grpc.OpenShellStub(channel)
|
|
with pytest.raises(grpc.RpcError) as exc_info:
|
|
stub.Health(openshell_pb2.HealthRequest(), timeout=10)
|
|
assert exc_info.value.code() in (
|
|
grpc.StatusCode.UNAVAILABLE,
|
|
grpc.StatusCode.UNKNOWN,
|
|
), f"expected UNAVAILABLE or UNKNOWN, got {exc_info.value.code()}"
|
|
finally:
|
|
channel.close()
|
|
|
|
def test_plaintext_connection_rejected(
|
|
self,
|
|
server_endpoint: tuple[str, int, str],
|
|
mtls_certs: tuple[bytes, bytes, bytes],
|
|
) -> None:
|
|
"""A plaintext (non-TLS) connection to the server port is rejected."""
|
|
host, port, _ = server_endpoint
|
|
# Ensure we have certs loaded (so the test isn't skipped for non-TLS).
|
|
_ = mtls_certs
|
|
|
|
channel = grpc.insecure_channel(f"{host}:{port}")
|
|
try:
|
|
stub = openshell_pb2_grpc.OpenShellStub(channel)
|
|
with pytest.raises(grpc.RpcError) as exc_info:
|
|
stub.Health(openshell_pb2.HealthRequest(), timeout=10)
|
|
# The loopback listener may intentionally accept plaintext service
|
|
# HTTP. A gRPC request is still rejected, either at the transport
|
|
# boundary or as an unimplemented HTTP route.
|
|
assert exc_info.value.code() in (
|
|
grpc.StatusCode.UNAVAILABLE,
|
|
grpc.StatusCode.UNKNOWN,
|
|
grpc.StatusCode.INTERNAL,
|
|
grpc.StatusCode.UNIMPLEMENTED,
|
|
), f"expected plaintext gRPC rejection, got {exc_info.value.code()}"
|
|
finally:
|
|
channel.close()
|