Files
OpenShell/e2e/python/test_security_tls.py
Drew Newberry 021400be8a refactor(auth): separate sandbox identity from TLS (#3110)
* 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>
2026-10-01 04:33:25 +00:00

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()