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