# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 from __future__ import annotations import json import os import pickle import threading import time from copy import deepcopy from dataclasses import asdict from pathlib import Path from types import SimpleNamespace from typing import Any, cast import pytest import openshell.sandbox as sandbox_module from openshell._proto import openshell_pb2 from openshell.mutations import DeletionOutcome from openshell.sandbox import ( _OIDC_TOKEN_EXPIRY_GRACE_SECONDS, _PYTHON_CLOUDPICKLE_BOOTSTRAP, _SANDBOX_PYTHON_BIN, ClientCredentialsAuth, Page, Pager, Sandbox, SandboxClient, SandboxError, SandboxRef, SandboxStatusRef, SandboxTemplateClient, ServiceAuthorizationMode, ServiceExposure, TlsConfig, _atomic_replace, _BearerAuthInterceptor, _load_cluster_bearer_token, _make_cluster_bearer_provider, _normalize_bearer, _OidcRefresher, _read_oidc_token_bundle, _sandbox_ref, _validate_oauth_url, ) def _request_workspace(request: Any) -> str | None: scope = request.workspace_scope if scope.WhichOneof("selection") == "workspace": return cast("str", scope.workspace) return None def _request_selects_all_workspaces(request: Any) -> bool: return request.workspace_scope.WhichOneof("selection") == "all_workspaces" def _request_sandbox(request: Any) -> str: name = getattr(request, "name", "") if name: return cast("str", name) return cast("str", request.sandbox) def _client_credentials_fixture() -> dict[str, Any]: return json.loads( ( Path(__file__).parents[2] / "sdk/conformance/oauth-client-credentials.json" ).read_text() ) def test_oauth_client_credentials_conformance_fixture() -> None: fixture = _client_credentials_fixture() assert fixture["expiry"]["leeway_seconds"] == _OIDC_TOKEN_EXPIRY_GRACE_SECONDS for value in fixture["urls"]["allowed"]: assert _validate_oauth_url("issuer", value) == value for value in fixture["urls"]["rejected"]: with pytest.raises(SandboxError): _validate_oauth_url("issuer", value) def test_client_credentials_auth_exact_form_cache_and_redaction() -> None: fixture = _client_credentials_fixture() seen: list[tuple[str, bytes]] = [] def handler(request: Any) -> Any: import httpx seen.append((str(request.url), bytes(request.content))) if request.method == "GET": return httpx.Response( 200, json={ "issuer": "https://issuer.example.com/", "token_endpoint": "https://issuer.example.com/token", }, ) return httpx.Response( 200, json={ "access_token": "service-token", "expires_in": fixture["expiry"]["valid_expires_in"], }, ) import httpx auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="service-client", client_secret="conformance-secret", scopes=("sandbox:read", "sandbox:write"), audience="openshell-gateway", _transport=httpx.MockTransport(handler), ) assert auth() == "service-token" assert auth() == "service-token" assert len(seen) == 2 form = dict(__import__("urllib.parse").parse.parse_qsl(seen[1][1].decode())) assert form == { field: fixture["request"][field] for field in ( "grant_type", "client_id", "client_secret", "scope", "audience", ) } assert "conformance-secret" not in repr(auth) def test_client_credentials_auth_preserves_explicit_empty_scopes() -> None: import httpx form: dict[str, str] = {} def handler(request: Any) -> Any: if request.method == "GET": return httpx.Response( 200, json={ "issuer": "https://issuer.example.com", "token_endpoint": "https://issuer.example.com/token", }, ) form.update( __import__("urllib.parse").parse.parse_qsl(request.content.decode()) ) return httpx.Response(200, json={"access_token": "token", "expires_in": 120}) auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="client", client_secret="secret", scopes=(), _transport=httpx.MockTransport(handler), ) auth._apply_gateway_metadata({"oidc_scopes": "sandbox:read sandbox:write"}) assert auth() == "token" assert auth._scopes == () assert "scope" not in form @pytest.mark.parametrize( "expires_in", _client_credentials_fixture()["expiry"]["invalid_expires_in"] ) def test_client_credentials_auth_rejects_invalid_expiry(expires_in: object) -> None: import httpx fixture = _client_credentials_fixture() def handler(request: Any) -> Any: if request.method == "GET": return httpx.Response( 200, json={ "issuer": fixture["discovery"]["matching_issuer"], "token_endpoint": "https://issuer.example.com/token", }, ) return httpx.Response( 200, json={"access_token": "token", "expires_in": expires_in} ) auth = ClientCredentialsAuth( issuer=fixture["discovery"]["configured_issuer"], client_id="client", client_secret="secret", _transport=httpx.MockTransport(handler), ) with pytest.raises(SandboxError, match="positive finite expires_in"): auth() @pytest.mark.parametrize( "status", _client_credentials_fixture()["discovery"]["redirect_statuses"] ) def test_client_credentials_auth_refuses_discovery_redirect(status: int) -> None: import httpx requests = 0 def handler(_request: Any) -> Any: nonlocal requests requests += 1 return httpx.Response( status, headers={"location": "https://attacker.example.com/discovery"}, ) auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="client", client_secret="secret", _transport=httpx.MockTransport(handler), ) with pytest.raises(SandboxError, match=f"HTTP {status}"): auth() assert requests == 1 @pytest.mark.parametrize( "status", _client_credentials_fixture()["discovery"]["redirect_statuses"] ) def test_client_credentials_auth_refuses_token_redirect(status: int) -> None: import httpx requests: list[str] = [] def handler(request: Any) -> Any: requests.append(str(request.url)) if request.method == "GET": return httpx.Response( 200, json={ "issuer": "https://issuer.example.com", "token_endpoint": "https://issuer.example.com/token", }, ) return httpx.Response( status, headers={"location": "https://attacker.example.com/token"}, ) auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="client", client_secret="secret", _transport=httpx.MockTransport(handler), ) with pytest.raises(SandboxError, match=f"HTTP {status}"): auth() assert requests == [ "https://issuer.example.com/.well-known/openid-configuration", "https://issuer.example.com/token", ] def test_client_credentials_auth_rejects_discovery_issuer_mismatch() -> None: import httpx fixture = _client_credentials_fixture() auth = ClientCredentialsAuth( issuer=fixture["discovery"]["configured_issuer"], client_id="client", client_secret="secret", _transport=httpx.MockTransport( lambda _request: httpx.Response( 200, json={ "issuer": fixture["discovery"]["mismatched_issuer"], "token_endpoint": "https://attacker.example.com/token", }, ) ), ) with pytest.raises(SandboxError, match="issuer mismatch"): auth() def test_client_credentials_auth_rejects_oversized_response() -> None: import httpx fixture = _client_credentials_fixture() auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="client", client_secret="secret", _transport=httpx.MockTransport( lambda _request: httpx.Response( 200, content=b"x" * (fixture["limits"]["max_response_bytes"] + 1) ) ), ) with pytest.raises(SandboxError, match="too large"): auth() def test_client_credentials_auth_single_flight_and_retry() -> None: import httpx token_calls = 0 release = threading.Event() def handler(request: Any) -> Any: nonlocal token_calls if request.method == "GET": return httpx.Response( 200, json={ "issuer": "http://127.0.0.1:8080", "token_endpoint": "http://127.0.0.1:8080/token", }, ) token_calls += 1 release.wait(timeout=2) return httpx.Response(200, json={"access_token": "shared", "expires_in": 120}) auth = ClientCredentialsAuth( issuer="http://127.0.0.1:8080", client_id="client", client_secret="secret", _transport=httpx.MockTransport(handler), ) results: list[str] = [] threads = [ threading.Thread(target=lambda: results.append(auth())) for _ in range(8) ] for thread in threads: thread.start() while token_calls == 0: time.sleep(0.001) release.set() for thread in threads: thread.join() assert results == ["shared"] * 8 assert token_calls == 1 def test_client_credentials_auth_fails_closed_and_redacts_errors() -> None: import httpx def supplier() -> str: raise RuntimeError("supplier-sensitive-detail") auth = ClientCredentialsAuth( issuer="http://localhost:8080", client_id="client", client_secret=supplier, _transport=httpx.MockTransport( lambda _request: httpx.Response( 200, json={ "issuer": "http://localhost:8080", "token_endpoint": "http://localhost:8080/token", }, ) ), ) with pytest.raises(SandboxError, match="supplier failed") as exc_info: auth() assert "supplier-sensitive-detail" not in str(exc_info.value) with pytest.raises(SandboxError, match="must use HTTPS"): ClientCredentialsAuth( issuer="http://remote.example.com", client_id="client", client_secret="secret", )() def test_client_credentials_auth_does_not_use_stale_token_after_renewal_failure() -> ( None ): import httpx exchanges = 0 def handler(request: Any) -> Any: nonlocal exchanges if request.method == "GET": return httpx.Response( 200, json={ "issuer": "http://localhost:8080", "token_endpoint": "http://localhost:8080/token", }, ) exchanges += 1 if exchanges == 1: return httpx.Response(200, json={"access_token": "stale", "expires_in": 30}) return httpx.Response(503, json={"error": "provider-sensitive-detail"}) auth = ClientCredentialsAuth( issuer="http://localhost:8080", client_id="client", client_secret="secret", _transport=httpx.MockTransport(handler), ) assert auth() == "stale" with pytest.raises(SandboxError, match="HTTP 503") as exc_info: auth() assert "stale" not in str(exc_info.value) assert "provider-sensitive-detail" not in str(exc_info.value) class _FakeStub: def __init__(self) -> None: self.request: openshell_pb2.ExecSandboxRequest | None = None def ExecSandbox( self, request: openshell_pb2.ExecSandboxRequest, timeout: float | None = None, ): self.request = request _ = timeout yield openshell_pb2.ExecSandboxEvent( exit=openshell_pb2.ExecSandboxExit(exit_code=0) ) def _client_with_fake_stub(stub: object) -> SandboxClient: client = cast("SandboxClient", object.__new__(SandboxClient)) client._timeout = 30.0 client._stub = cast("Any", stub) return client def _template_client_with_fake_stub(stub: object) -> SandboxTemplateClient: client = cast("SandboxTemplateClient", object.__new__(SandboxTemplateClient)) client._timeout = 30.0 client._stub = cast("Any", stub) return client def test_exec_sends_stdin_payload() -> None: stub = _FakeStub() client = _client_with_fake_stub(stub) result = client.exec( "sandbox-1", ["python", "-c", "print('ok')"], workspace="default", stdin=b"payload", ) assert result.exit_code == 0 assert stub.request is not None assert stub.request.stdin == b"payload" def test_exec_python_serializes_callable_payload() -> None: stub = _FakeStub() client = _client_with_fake_stub(stub) def add(a: int, b: int) -> int: return a + b result = client.exec_python("sandbox-1", add, workspace="default", args=(2, 3)) assert result.exit_code == 0 assert stub.request is not None assert stub.request.command == [ _SANDBOX_PYTHON_BIN, "-c", _PYTHON_CLOUDPICKLE_BOOTSTRAP, ] assert stub.request.environment["OPENSHELL_PYFUNC_B64"] assert stub.request.stdin == b"" def test_from_active_cluster_reads_gateway_metadata_layout( tmp_path: Path, monkeypatch: Any, ) -> None: gateway_name = "test-gateway" gateway_dir = tmp_path / "openshell" / "gateways" / gateway_name mtls_dir = gateway_dir / "mtls" mtls_dir.mkdir(parents=True) (tmp_path / "openshell" / "active_gateway").write_text(gateway_name) (gateway_dir / "metadata.json").write_text( json.dumps({"gateway_endpoint": "https://127.0.0.1:8443"}) ) (mtls_dir / "ca.crt").write_text("ca") (mtls_dir / "tls.crt").write_text("cert") (mtls_dir / "tls.key").write_text("key") monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) monkeypatch.delenv("OPENSHELL_GATEWAY", raising=False) client = SandboxClient.from_active_cluster() try: assert client._cluster_name == gateway_name finally: client.close() def test_from_active_cluster_prefers_openshell_gateway_env( tmp_path: Path, monkeypatch: Any, ) -> None: gateway_name = "env-gateway" gateway_dir = tmp_path / "openshell" / "gateways" / gateway_name mtls_dir = gateway_dir / "mtls" mtls_dir.mkdir(parents=True) (gateway_dir / "metadata.json").write_text( json.dumps({"gateway_endpoint": "https://127.0.0.1:8443"}) ) (mtls_dir / "ca.crt").write_text("ca") (mtls_dir / "tls.crt").write_text("cert") (mtls_dir / "tls.key").write_text("key") monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) monkeypatch.setenv("OPENSHELL_GATEWAY", gateway_name) client = SandboxClient.from_active_cluster() try: assert client._cluster_name == gateway_name finally: client.close() # --------------------------------------------------------------------------- # OIDC bearer auth # --------------------------------------------------------------------------- class _FakeClientCallDetails: """grpc.ClientCallDetails is a NamedTuple in real gRPC; for unit tests we just need an object with the same field set and a ._replace shim.""" __slots__ = ("credentials", "metadata", "method", "timeout", "wait_for_ready") def __init__( self, method: str = "/Test/Method", timeout: float | None = None, metadata: Any = None, credentials: Any = None, wait_for_ready: Any = None, ) -> None: self.method = method self.timeout = timeout self.metadata = metadata self.credentials = credentials self.wait_for_ready = wait_for_ready def _replace(self, **kwargs: Any) -> _FakeClientCallDetails: return _FakeClientCallDetails( method=kwargs.get("method", self.method), timeout=kwargs.get("timeout", self.timeout), metadata=kwargs.get("metadata", self.metadata), credentials=kwargs.get("credentials", self.credentials), wait_for_ready=kwargs.get("wait_for_ready", self.wait_for_ready), ) def test_normalize_bearer_accepts_str_or_callable() -> None: assert _normalize_bearer(None) is None static = _normalize_bearer("abc") assert static is not None assert static() == "abc" counter = [0] def provider() -> str: counter[0] += 1 return f"token-{counter[0]}" dynamic = _normalize_bearer(provider) assert dynamic is not None assert dynamic() == "token-1" assert dynamic() == "token-2" def test_bearer_interceptor_attaches_authorization_header() -> None: interceptor = _BearerAuthInterceptor(lambda: "secret-token") captured: dict[str, Any] = {} def continuation(details: Any, request: Any) -> str: captured["details"] = details captured["request"] = request return "result" details = _FakeClientCallDetails(metadata=[("x-existing", "yes")]) result = interceptor.intercept_unary_unary(continuation, details, "payload") assert result == "result" md = list(captured["details"].metadata) # Pre-existing metadata preserved, authorization appended last. assert ("x-existing", "yes") in md assert ("authorization", "Bearer secret-token") in md assert captured["request"] == "payload" def test_bearer_interceptor_handles_empty_metadata() -> None: interceptor = _BearerAuthInterceptor(lambda: "t") captured: dict[str, Any] = {} def continuation(details: Any, _request: Any) -> None: captured["metadata"] = list(details.metadata) details = _FakeClientCallDetails(metadata=None) interceptor.intercept_unary_unary(continuation, details, request="x") assert captured["metadata"] == [("authorization", "Bearer t")] def test_bearer_interceptor_calls_token_provider_per_request() -> None: tokens = iter(["t1", "t2", "t3"]) interceptor = _BearerAuthInterceptor(lambda: next(tokens)) seen: list[str] = [] def continuation(details: Any, _request: Any) -> None: for key, value in details.metadata: if key == "authorization": seen.append(value) for _ in range(3): interceptor.intercept_unary_unary( continuation, _FakeClientCallDetails(), request="x" ) assert seen == ["Bearer t1", "Bearer t2", "Bearer t3"] def test_load_cluster_bearer_token_reads_oidc_token_json(tmp_path: Path) -> None: gateway_dir = tmp_path / "gw" gateway_dir.mkdir() (gateway_dir / "oidc_token.json").write_text( json.dumps( { "access_token": "jwt-blob", "refresh_token": "rt", "expires_at": 9999999999, "issuer": "https://idp.example/realms/openshell", "client_id": "openshell-cli", } ) ) assert _load_cluster_bearer_token(gateway_dir) == "jwt-blob" def test_load_cluster_bearer_token_returns_none_when_missing( tmp_path: Path, ) -> None: assert _load_cluster_bearer_token(tmp_path / "absent") is None def test_load_cluster_bearer_token_tolerates_unreadable_file( tmp_path: Path, ) -> None: gateway_dir = tmp_path / "gw" gateway_dir.mkdir() (gateway_dir / "oidc_token.json").write_text("not json") assert _load_cluster_bearer_token(gateway_dir) is None def test_load_cluster_bearer_token_rejects_missing_access_token( tmp_path: Path, ) -> None: gateway_dir = tmp_path / "gw" gateway_dir.mkdir() (gateway_dir / "oidc_token.json").write_text(json.dumps({"refresh_token": "rt"})) assert _load_cluster_bearer_token(gateway_dir) is None def _setup_gateway_dir( tmp_path: Path, monkeypatch: Any, *, name: str = "g", endpoint: str = "http://127.0.0.1:8080", auth_mode: str | None = None, mtls_files: dict[str, str] | None = None, oidc_bundle: dict | None = None, ) -> Path: gateway_dir = tmp_path / "openshell" / "gateways" / name gateway_dir.mkdir(parents=True) (tmp_path / "openshell" / "active_gateway").write_text(name) meta: dict[str, Any] = {"gateway_endpoint": endpoint} if auth_mode is not None: meta["auth_mode"] = auth_mode (gateway_dir / "metadata.json").write_text(json.dumps(meta)) if mtls_files: mtls_dir = gateway_dir / "mtls" mtls_dir.mkdir() for fname, body in mtls_files.items(): (mtls_dir / fname).write_text(body) if oidc_bundle is not None: (gateway_dir / "oidc_token.json").write_text(json.dumps(oidc_bundle)) monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) monkeypatch.delenv("OPENSHELL_GATEWAY", raising=False) return gateway_dir def _channel_is_intercepted(channel: Any) -> bool: """grpc.intercept_channel returns a _Channel whose module name ends in `interceptor`. We don't depend on the class name (it varies across gRPC versions); module is stable.""" return type(channel).__module__.endswith("interceptor") def test_from_active_cluster_loads_bearer_when_auth_mode_is_oidc( tmp_path: Path, monkeypatch: Any, ) -> None: """Finding 3: bearer is attached iff metadata.auth_mode == "oidc".""" _setup_gateway_dir( tmp_path, monkeypatch, auth_mode="oidc", oidc_bundle={"access_token": "from-disk"}, ) client = SandboxClient.from_active_cluster() try: assert _channel_is_intercepted(client._channel) finally: client.close() def test_from_active_cluster_ignores_stale_token_when_auth_mode_not_oidc( tmp_path: Path, monkeypatch: Any, ) -> None: """Finding 3: a stale oidc_token.json alongside a non-OIDC gateway must NOT cause bearer auth to be attached.""" _setup_gateway_dir( tmp_path, monkeypatch, # auth_mode omitted (or "mtls", "plaintext") — anything but "oidc". oidc_bundle={"access_token": "stale-from-disk"}, ) client = SandboxClient.from_active_cluster() try: # Plain channel, no interceptor wrapper. assert not _channel_is_intercepted(client._channel) finally: client.close() def test_from_active_cluster_https_oidc_without_mtls_uses_tls_with_system_roots( tmp_path: Path, monkeypatch: Any, ) -> None: """Finding 1: https OIDC gateways without mTLS material must still use a TLS channel (system roots) — NOT fall back to insecure_channel.""" _setup_gateway_dir( tmp_path, monkeypatch, endpoint="https://gateway.example:443", auth_mode="oidc", oidc_bundle={"access_token": "t"}, ) client = SandboxClient.from_active_cluster() try: # The bearer interceptor wraps the channel, so inspect the # wrapped channel's class to confirm it's a secure (TLS) channel. inner = getattr(client._channel, "_channel", client._channel) # gRPC's `grpc.secure_channel` returns a `_Channel` from # `grpc._channel`; we can't trivially introspect "secure" vs # "insecure" on the wrapper itself. Probe by attempting to # extract the connectivity state — both kinds expose it — and # rely on a behavioral assertion: an insecure channel against # a hostname-only endpoint would have already attached TCP-only # subchannels. Easier: verify TlsConfig() was used by checking # the SandboxClient endpoint normalized correctly. # The most direct assertion is on the client config: assert client._endpoint == "gateway.example:443" # And the channel must not be insecure. assert "InsecureChannelCredentials" not in repr(inner) finally: client.close() def test_from_active_cluster_https_ca_only_layout( tmp_path: Path, monkeypatch: Any, ) -> None: """Finding 1: a CA-only mtls directory (ca.crt but no tls.crt/tls.key) must produce a CA-only TLS channel, not a FileNotFoundError.""" _setup_gateway_dir( tmp_path, monkeypatch, endpoint="https://gateway.example:443", auth_mode="oidc", mtls_files={ "ca.crt": "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----\n" }, oidc_bundle={"access_token": "t"}, ) # Should not raise. client = SandboxClient.from_active_cluster() try: assert client._endpoint == "gateway.example:443" finally: client.close() def test_tls_config_rejects_partial_client_identity() -> None: """Cert without key (or vice versa) is a misconfiguration.""" import pytest as _pytest with _pytest.raises(ValueError, match="cert_path and key_path"): TlsConfig(cert_path=Path("/x.crt")) def test_tls_config_allows_empty_for_system_roots() -> None: """`TlsConfig()` is the system-roots-trust flavor.""" cfg = TlsConfig() assert cfg.ca_path is None and cfg.cert_path is None and cfg.key_path is None # --------------------------------------------------------------------------- # Provider semantics: per-RPC reload + expiry # --------------------------------------------------------------------------- def test_cluster_bearer_provider_reloads_on_every_call(tmp_path: Path) -> None: """The fail-closed (no-refresh) provider re-reads oidc_token.json each invocation, so a long-lived SandboxClient picks up CLI rotations without reconstruction.""" gateway_dir = tmp_path token_file = gateway_dir / "oidc_token.json" token_file.write_text(json.dumps({"access_token": "first"})) provider, _ = _make_cluster_bearer_provider(gateway_dir, "g", auto_refresh=False) assert provider() == "first" # Simulate `openshell gateway login` writing a new token. token_file.write_text(json.dumps({"access_token": "second"})) assert provider() == "second" def test_cluster_bearer_provider_raises_on_expired_token(tmp_path: Path) -> None: """Fail-closed provider raises on expiry with a clear re-login hint.""" gateway_dir = tmp_path (gateway_dir / "oidc_token.json").write_text( json.dumps({"access_token": "expired", "expires_at": 1}) ) provider, _ = _make_cluster_bearer_provider( gateway_dir, "stale-gateway", auto_refresh=False ) import pytest as _pytest with _pytest.raises(SandboxError, match="expired"): provider() def test_cluster_bearer_provider_raises_when_file_missing(tmp_path: Path) -> None: provider, _ = _make_cluster_bearer_provider( tmp_path / "absent", "g", auto_refresh=False ) import pytest as _pytest with _pytest.raises(SandboxError, match="missing or unreadable"): provider() def test_cluster_bearer_provider_raises_on_missing_access_token( tmp_path: Path, ) -> None: (tmp_path / "oidc_token.json").write_text(json.dumps({"refresh_token": "r"})) provider, _ = _make_cluster_bearer_provider(tmp_path, "g", auto_refresh=False) import pytest as _pytest with _pytest.raises(SandboxError, match="no access token"): provider() # --------------------------------------------------------------------------- # OAuth2 native refresh (_OidcRefresher) — opt-in via auto_refresh=True. # --------------------------------------------------------------------------- def _write_bundle( gateway_dir: Path, *, access_token: str = "fresh", refresh_token: str = "r-orig", expires_at: int | None = None, issuer: str = "https://idp.example/realms/openshell", client_id: str = "openshell-cli", ) -> None: bundle: dict[str, Any] = { "access_token": access_token, "refresh_token": refresh_token, "issuer": issuer, "client_id": client_id, } if expires_at is not None: bundle["expires_at"] = expires_at (gateway_dir / "oidc_token.json").write_text(json.dumps(bundle)) DEFAULT_ISSUER = "https://idp.example/realms/openshell" DEFAULT_TOKEN_ENDPOINT = ( "https://idp.example/realms/openshell/protocol/openid-connect/token" ) def _make_mock_transport( *, discovery: dict | None = None, refresh_responses: list[dict] | None = None, discovery_status: int = 200, refresh_status: int = 200, seen_refresh: list[Any] | None = None, seen_discovery: list[Any] | None = None, ): """Build an httpx.MockTransport that serves OIDC discovery + token refresh from an in-memory script. `refresh_responses` is consumed in order across successive POSTs to the token endpoint (which lets tests assert refresh-token rotation semantics across multiple refreshes). """ import httpx as _httpx refresh_iter = iter( refresh_responses or [{"access_token": "refreshed-jwt", "expires_in": 3600}] ) def handler(request: _httpx.Request) -> _httpx.Response: if request.url.path.endswith("/.well-known/openid-configuration"): if seen_discovery is not None: seen_discovery.append(str(request.url)) body = discovery or { "issuer": DEFAULT_ISSUER, "token_endpoint": DEFAULT_TOKEN_ENDPOINT, } return _httpx.Response(discovery_status, json=body) # Anything else is a refresh exchange. if seen_refresh is not None: seen_refresh.append((str(request.url), bytes(request.content))) try: body = next(refresh_iter) except StopIteration: return _httpx.Response(500, json={"error": "test_script_exhausted"}) return _httpx.Response(refresh_status, json=body) return _httpx.MockTransport(handler) def _install_mock_transport(refresher: Any, transport: Any) -> None: """Swap the refresher's httpx.Client for one bound to a mock transport. We rebuild with `follow_redirects=False` so the redirect-rejection test still exercises the real policy. """ import httpx as _httpx refresher._http.close() refresher._http = _httpx.Client(transport=transport, follow_redirects=False) def test_refresher_returns_cached_token_when_fresh(tmp_path: Path) -> None: """No refresh round-trip when the cached bundle is still fresh.""" _write_bundle(tmp_path, expires_at=int(time.time()) + 3600) seen: list[Any] = [] transport = _make_mock_transport( seen_discovery=seen, seen_refresh=seen, ) r = _OidcRefresher(tmp_path, "g") _install_mock_transport(r, transport) assert r.current_access_token() == "fresh" assert seen == [] # no discovery, no refresh def test_refresher_picks_up_disk_rotation_before_refreshing( tmp_path: Path, ) -> None: """If the in-memory bundle is stale but the CLI just wrote a fresh one, re-read disk instead of hitting the IdP.""" _write_bundle(tmp_path, access_token="old", expires_at=1) seen_refresh: list[Any] = [] transport = _make_mock_transport(seen_refresh=seen_refresh) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) # First call: refresh against IdP — exercise that path first. assert r.current_access_token() == "refreshed-jwt" assert len(seen_refresh) == 1 # Now simulate the CLI writing a fresh bundle. Force the in-memory # state to look stale so the disk re-read path triggers. _write_bundle( tmp_path, access_token="cli-rotated", expires_at=int(time.time()) + 3600 ) r._bundle = { "access_token": "stale-in-memory", "expires_at": 1, "refresh_token": "r", } # Replace the transport with one that asserts on any request. import httpx as _httpx def assert_no_calls(_req: _httpx.Request) -> _httpx.Response: raise AssertionError("should not refresh — disk was fresh") _install_mock_transport(r, _httpx.MockTransport(assert_no_calls)) assert r.current_access_token() == "cli-rotated" def test_refresher_adopts_stale_disk_refresh_token_before_refreshing( tmp_path: Path, ) -> None: """Regression: when both the in-memory and on-disk access tokens are stale but another process rotated the on-disk refresh_token, refresh with the disk refresh_token (r2), not the invalidated in-memory one (r1). Without this, a rotating IdP (Keycloak with rotation, Entra strict) would invalid_grant because process A still holds the pre-rotation r1. """ # Disk holds a rotated-but-stale bundle (r2) written by another process. # Its access token was minted more recently than ours (later expiry, # though still inside the grace window), so disk carries the newer # refresh_token even though both are due for refresh. disk_exp = int(time.time()) + 5 _write_bundle( tmp_path, access_token="disk-old", expires_at=disk_exp, refresh_token="r2" ) seen: list[Any] = [] transport = _make_mock_transport( refresh_responses=[ {"access_token": "a-new", "refresh_token": "r3", "expires_in": 3600}, ], seen_refresh=seen, ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) # Seed older stale in-memory state holding the pre-rotation token r1. r._bundle = { "access_token": "mem-old", "expires_at": 1, "refresh_token": "r1", "issuer": DEFAULT_ISSUER, } assert r.current_access_token() == "a-new" # The refresh POST must carry the disk's r2, never the stale r1. _, body = seen[-1] assert b"refresh_token=r2" in body assert b"refresh_token=r1" not in body def test_refresher_resets_token_endpoint_when_disk_issuer_changes( tmp_path: Path, ) -> None: """When the adopted disk bundle has a different issuer than the cached one, the previously discovered token endpoint must be re-discovered against the new issuer rather than reused.""" new_issuer = "https://other-idp.example/realms/openshell" # Disk is newer than the in-memory bundle (later expiry) so it is # adopted, but still stale so a refresh — and thus re-discovery — runs. _write_bundle( tmp_path, access_token="disk-old", expires_at=int(time.time()) + 5, refresh_token="r2", issuer=new_issuer, ) seen_discovery: list[Any] = [] transport = _make_mock_transport( discovery={ "issuer": new_issuer, "token_endpoint": f"{new_issuer}/protocol/openid-connect/token", }, seen_discovery=seen_discovery, ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) # Pretend we already discovered an endpoint for the OLD issuer. r._token_endpoint = f"{DEFAULT_ISSUER}/protocol/openid-connect/token" r._bundle = { "access_token": "mem-old", "expires_at": 1, "refresh_token": "r1", "issuer": DEFAULT_ISSUER, } r.current_access_token() # Re-discovery happened against the new issuer. assert len(seen_discovery) == 1 assert new_issuer in seen_discovery[0] def test_refresher_recovers_from_invalid_grant_after_peer_rotation( tmp_path: Path, ) -> None: """If our refresh POST loses a rotation race (peer already rotated r1→r2 and the IdP rejects our r1 with invalid_grant), re-read disk, pick up the peer's r2, and retry — succeeding without forcing a re-authenticate.""" import httpx as _httpx _write_bundle(tmp_path, access_token="old", expires_at=1, refresh_token="r1") posts: list[bytes] = [] def handler(request: _httpx.Request) -> _httpx.Response: if request.url.path.endswith("/.well-known/openid-configuration"): return _httpx.Response( 200, json={ "issuer": DEFAULT_ISSUER, "token_endpoint": DEFAULT_TOKEN_ENDPOINT, }, ) body = bytes(request.content) posts.append(body) if b"refresh_token=r1" in body: # Simulate the peer: it already rotated r1→r2 and wrote r2 to # disk, so the IdP rejects our now-stale r1. _write_bundle( tmp_path, access_token="peer", expires_at=int(time.time()) + 5, refresh_token="r2", ) return _httpx.Response(400, json={"error": "invalid_grant"}) # The retry carries the peer's r2 and succeeds. return _httpx.Response( 200, json={"access_token": "a-final", "refresh_token": "r3", "expires_in": 3600}, ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, _httpx.MockTransport(handler)) assert r.current_access_token() == "a-final" # Exactly two refresh POSTs: the failed r1 then the recovered r2. assert any(b"refresh_token=r1" in p for p in posts) assert any(b"refresh_token=r2" in p for p in posts) assert len(posts) == 2 def test_refresher_reraises_invalid_grant_without_peer_rotation( tmp_path: Path, ) -> None: """invalid_grant with no peer rotation (disk still holds our refresh_token) is a genuine dead token — surface the re-authenticate hint and do NOT loop on the retry path.""" import httpx as _httpx import pytest as _pytest _write_bundle(tmp_path, access_token="old", expires_at=1, refresh_token="r1") posts: list[bytes] = [] def handler(request: _httpx.Request) -> _httpx.Response: if request.url.path.endswith("/.well-known/openid-configuration"): return _httpx.Response( 200, json={ "issuer": DEFAULT_ISSUER, "token_endpoint": DEFAULT_TOKEN_ENDPOINT, }, ) posts.append(bytes(request.content)) return _httpx.Response(400, json={"error": "invalid_grant"}) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, _httpx.MockTransport(handler)) with _pytest.raises(SandboxError, match="Re-authenticate"): r.current_access_token() # Only one POST — disk offered no new refresh_token, so no retry. assert len(posts) == 1 def test_refresher_exchanges_refresh_token_when_stale(tmp_path: Path) -> None: """When both memory and disk are stale, do the OAuth2 refresh exchange.""" _write_bundle(tmp_path, access_token="old", expires_at=1) seen_refresh: list[Any] = [] transport = _make_mock_transport(seen_refresh=seen_refresh) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) assert r.current_access_token() == "refreshed-jwt" # The refresh request should be a POST to the discovered token endpoint # with grant_type=refresh_token in the body. url, body = seen_refresh[-1] assert url.endswith("/protocol/openid-connect/token") assert b"grant_type=refresh_token" in body assert b"refresh_token=r-orig" in body def test_refresher_writes_back_when_enabled(tmp_path: Path) -> None: """write_back=True persists rotated bundle to disk atomically with 0600.""" _write_bundle(tmp_path, access_token="old", expires_at=1) transport = _make_mock_transport( refresh_responses=[ { "access_token": "rotated", "refresh_token": "r-new", "expires_in": 3600, } ], ) r = _OidcRefresher(tmp_path, "g", write_back=True) _install_mock_transport(r, transport) assert r.current_access_token() == "rotated" on_disk = json.loads((tmp_path / "oidc_token.json").read_text()) assert on_disk["access_token"] == "rotated" assert on_disk["refresh_token"] == "r-new" # Mode should be 0600 on POSIX. if os.name == "posix": mode = (tmp_path / "oidc_token.json").stat().st_mode & 0o777 assert mode == 0o600, f"got {oct(mode)}" def test_refresher_write_back_is_default(tmp_path: Path) -> None: """Default IS write_back=True so refresh-token rotation propagates to disk for other processes (Rust CLI, TUI, second Python client).""" _write_bundle(tmp_path, access_token="old", expires_at=1) transport = _make_mock_transport( refresh_responses=[ { "access_token": "rotated", "refresh_token": "r-new", "expires_in": 3600, } ], ) r = _OidcRefresher(tmp_path, "g") # default write_back=True _install_mock_transport(r, transport) r.current_access_token() on_disk = json.loads((tmp_path / "oidc_token.json").read_text()) assert on_disk["access_token"] == "rotated" assert on_disk["refresh_token"] == "r-new" def test_refresher_honors_refresh_token_rotation(tmp_path: Path) -> None: """When the IdP returns a new refresh_token, use it for subsequent refreshes instead of the original. Some IdPs (Keycloak with rotation enabled, Entra in strict mode) reissue and invalidate the old refresh_token.""" _write_bundle(tmp_path, access_token="old", expires_at=1, refresh_token="r1") seen: list[Any] = [] transport = _make_mock_transport( refresh_responses=[ {"access_token": "a2", "refresh_token": "r2", "expires_in": 1}, {"access_token": "a3", "refresh_token": "r3", "expires_in": 3600}, ], seen_refresh=seen, ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) assert r.current_access_token() == "a2" # Second call: a2 is also expired (expires_in=1), so we refresh again, # this time the request body should carry the rotated r2 (not r1). assert r.current_access_token() == "a3" assert b"refresh_token=r1" in seen[0][1] assert b"refresh_token=r2" in seen[1][1] def test_refresher_second_process_can_refresh_after_rotation( tmp_path: Path, ) -> None: """Two-process simulation (Finding #2): process A refreshes r1→r2 with write_back=True (default). Process B starts from disk and successfully uses r2 — proving the rotation reached the shared cache.""" _write_bundle(tmp_path, access_token="old", expires_at=1, refresh_token="r1") transport_a = _make_mock_transport( refresh_responses=[ {"access_token": "a2", "refresh_token": "r2", "expires_in": 1}, ], ) process_a = _OidcRefresher(tmp_path, "g") # write_back=True (default) _install_mock_transport(process_a, transport_a) assert process_a.current_access_token() == "a2" # Process B picks up the cache fresh. The IdP now expects r2; if the # disk still held r1, this would fail at the IdP. With write_back the # disk has r2, and B refreshes successfully. seen_b: list[Any] = [] transport_b = _make_mock_transport( refresh_responses=[ {"access_token": "a3", "refresh_token": "r3", "expires_in": 3600}, ], seen_refresh=seen_b, ) process_b = _OidcRefresher(tmp_path, "g") _install_mock_transport(process_b, transport_b) assert process_b.current_access_token() == "a3" # Process B should have presented r2, not r1. assert b"refresh_token=r2" in seen_b[0][1] def test_refresher_concurrent_calls_share_one_refresh(tmp_path: Path) -> None: """N threads racing on a stale token should produce exactly one refresh exchange (not N). Mirrors google-auth's RefreshThreadManager coordination.""" _write_bundle(tmp_path, access_token="old", expires_at=1) refresh_count = [0] barrier = threading.Barrier(8) import httpx as _httpx def handler(request: _httpx.Request) -> _httpx.Response: if request.url.path.endswith("/.well-known/openid-configuration"): return _httpx.Response( 200, json={ "issuer": DEFAULT_ISSUER, "token_endpoint": DEFAULT_TOKEN_ENDPOINT, }, ) refresh_count[0] += 1 return _httpx.Response( 200, json={ "access_token": "refreshed", "expires_in": 3600, }, ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, _httpx.MockTransport(handler)) results: list[str] = [] errors: list[BaseException] = [] def worker() -> None: try: barrier.wait() results.append(r.current_access_token()) except BaseException as e: errors.append(e) threads = [threading.Thread(target=worker) for _ in range(8)] for t in threads: t.start() for t in threads: t.join() assert not errors, errors assert results == ["refreshed"] * 8 # One refresh exchange, regardless of thread count. assert refresh_count[0] == 1, f"expected one refresh, got {refresh_count[0]}" def test_refresher_surfaces_idp_failure_as_sandbox_error( tmp_path: Path, ) -> None: """A non-2xx from the token endpoint becomes a SandboxError.""" _write_bundle(tmp_path, access_token="old", expires_at=1) transport = _make_mock_transport( refresh_status=400, refresh_responses=[ { "error": "invalid_grant", "error_description": "Token is not active", } ], ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) import pytest as _pytest with _pytest.raises(SandboxError, match="refresh failed"): r.current_access_token() def test_refresher_rejects_issuer_mismatch_in_discovery(tmp_path: Path) -> None: """Finding #1 (Critical): a discovery doc claiming a different issuer must be rejected. Without this, a malicious or misdirected discovery response could steer the refresh_token POST to an attacker- controlled endpoint.""" _write_bundle(tmp_path, access_token="old", expires_at=1) transport = _make_mock_transport( discovery={ "issuer": "https://evil.example/realms/openshell", "token_endpoint": "https://evil.example/token", }, ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) import pytest as _pytest with _pytest.raises(SandboxError, match="issuer mismatch"): r.current_access_token() def test_refresher_rejects_redirect_during_discovery(tmp_path: Path) -> None: """Finding #1 (Critical): a 3xx during OIDC discovery must NOT be auto-followed — that would let a network attacker steer the SDK to an arbitrary token_endpoint URL. The Rust CLI sets `Policy::none()`; we set httpx's `follow_redirects=False`.""" _write_bundle(tmp_path, access_token="old", expires_at=1) transport = _make_mock_transport( discovery_status=302, discovery={"location": "https://evil.example/...."}, ) r = _OidcRefresher(tmp_path, "g", write_back=False) _install_mock_transport(r, transport) import pytest as _pytest with _pytest.raises(SandboxError, match=r"discovery failed.*HTTP 302"): r.current_access_token() def test_refresher_insecure_disables_tls_verification() -> None: """Finding #3: insecure=True propagates to httpx as verify=False so self-signed OIDC issuers work the same way they do in the Rust CLI's `--insecure` plumbing.""" import pathlib r = _OidcRefresher( pathlib.Path("/tmp/does-not-exist"), "g", insecure=True, ) try: # httpx exposes the configured verify policy on the client; we # don't depend on its precise type, just on it being a falsy # value (the default is True / an SSLContext). # In recent httpx versions this lives on the underlying transport. # The simplest stable check is: an insecure client allows # connect to self-signed hosts; the rest of the contract is # httpx's responsibility. # Verify the client's verify attribute (whether top-level or via # transport) is False. assert _client_verify_is_disabled(r._http) finally: r.close() def _client_verify_is_disabled(client: Any) -> bool: """Inspect an httpx.Client for verify=False. httpx surfaces verify either on the client directly (older) or via the default transport (newer).""" if getattr(client, "verify", None) is False: return True transport = getattr(client, "_transport", None) if transport is None: return False # httpx's default HTTPTransport wraps an SSL context or a bool. pool = getattr(transport, "_pool", None) if pool is not None: ssl_context = getattr(pool, "_ssl_context", None) # When verify=False, httpx builds a context without verification. if ssl_context is not None: import ssl return ssl_context.verify_mode == ssl.CERT_NONE # Fallback: check for any internal `_verify` attribute set to False. return getattr(transport, "_verify", None) is False def test_refresher_raises_when_bundle_has_no_refresh_token( tmp_path: Path, ) -> None: """Without a refresh_token (e.g. client_credentials grant — different code path entirely), refresh has nothing to exchange and surfaces a clear error.""" (tmp_path / "oidc_token.json").write_text( json.dumps({"access_token": "old", "expires_at": 1, "issuer": "x"}) ) r = _OidcRefresher(tmp_path, "g", write_back=False) import pytest as _pytest with _pytest.raises(SandboxError, match="no refresh token"): r.current_access_token() # --------------------------------------------------------------------------- # auth_mode gate: only metadata.json["auth_mode"] == "oidc" wires the bearer # interceptor. A stray oidc_token.json next to a non-OIDC gateway must not # trigger it. # --------------------------------------------------------------------------- def test_mtls_only_from_active_cluster_skips_bearer_interceptor( tmp_path: Path, monkeypatch: Any, ) -> None: """from_active_cluster against an mTLS-only gateway (no auth_mode set) does not wrap the channel with a bearer interceptor, even if a stale oidc_token.json is present in the gateway directory.""" gateway_name = "mtls-only" gateway_dir = tmp_path / "openshell" / "gateways" / gateway_name mtls_dir = gateway_dir / "mtls" mtls_dir.mkdir(parents=True) (tmp_path / "openshell" / "active_gateway").write_text(gateway_name) # No auth_mode field — the chart-default path. (gateway_dir / "metadata.json").write_text( json.dumps({"gateway_endpoint": "https://127.0.0.1:8443"}) ) for f in ("ca.crt", "tls.crt", "tls.key"): (mtls_dir / f).write_text(f"-----BEGIN {f}-----\n-----END {f}-----\n") # Stray oidc_token.json — proving the auth_mode gate (and not the # file's presence) is what would trigger the refresher. (gateway_dir / "oidc_token.json").write_text(json.dumps({"access_token": "stale"})) monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) monkeypatch.delenv("OPENSHELL_GATEWAY", raising=False) client = SandboxClient.from_active_cluster() try: # No bearer interceptor wraps the channel. assert not type(client._channel).__module__.endswith("interceptor") finally: client.close() # --------------------------------------------------------------------------- # Lifecycle plumbing: close() releases refresher resources, concurrent # write-back doesn't trample. # --------------------------------------------------------------------------- def test_sandbox_client_close_invokes_bearer_close() -> None: """`SandboxClient.close()` must invoke the `_bearer_close` callable wired by `from_active_cluster`. Otherwise the refresher's httpx.Client leaks sockets/FDs until GC runs.""" closed = [0] def bearer_close() -> None: closed[0] += 1 client = SandboxClient( "localhost:8080", bearer_token="tok", _bearer_close=bearer_close, ) client.close() assert closed[0] == 1 # close() is idempotent — re-invoking does not double-call. client.close() assert closed[0] == 1 def test_from_active_cluster_fills_client_credentials_from_metadata( tmp_path: Path, monkeypatch: Any, ) -> None: gateway_dir = _setup_gateway_dir(tmp_path, monkeypatch, auth_mode="oidc") metadata_path = gateway_dir / "metadata.json" metadata = json.loads(metadata_path.read_text()) metadata.update( { "oidc_issuer": "https://issuer.example.com", "oidc_client_id": "service-client", "oidc_audience": "gateway", "oidc_scopes": "sandbox:read sandbox:write", } ) metadata_path.write_text(json.dumps(metadata)) auth = ClientCredentialsAuth(client_secret="secret") client = SandboxClient.from_active_cluster( client_credentials=auth, insecure=True, ) try: assert auth._issuer == "https://issuer.example.com" assert auth._client_id == "service-client" assert auth._audience == "gateway" assert auth._scopes == ("sandbox:read", "sandbox:write") assert auth._insecure is True finally: client.close() def test_sandbox_client_rejects_client_credentials_on_remote_plaintext( monkeypatch: Any, ) -> None: channel_opened = False def insecure_channel(_endpoint: str) -> Any: nonlocal channel_opened channel_opened = True raise AssertionError("plaintext channel must not be opened") monkeypatch.setattr(sandbox_module.grpc, "insecure_channel", insecure_channel) auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="client", client_secret="secret", ) with pytest.raises(SandboxError, match="require TLS"): SandboxClient("gateway.example.com:50051", client_credentials=auth) assert not channel_opened @pytest.mark.parametrize( "endpoint", ["localhost:50051", "127.42.0.1:50051", "[::1]:50051"], ) def test_sandbox_client_allows_client_credentials_on_plaintext_loopback( endpoint: str, ) -> None: auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="client", client_secret="secret", ) client = SandboxClient(endpoint, client_credentials=auth) client.close() def test_from_active_cluster_rejects_client_credentials_on_remote_plaintext( tmp_path: Path, monkeypatch: Any, ) -> None: gateway_dir = _setup_gateway_dir( tmp_path, monkeypatch, endpoint="http://gateway.example.com:8080", auth_mode="oidc", ) metadata_path = gateway_dir / "metadata.json" metadata = json.loads(metadata_path.read_text()) metadata.update( { "oidc_issuer": "https://issuer.example.com", "oidc_client_id": "service-client", } ) metadata_path.write_text(json.dumps(metadata)) auth = ClientCredentialsAuth(client_secret="secret") with pytest.raises(SandboxError, match="require TLS"): SandboxClient.from_active_cluster(client_credentials=auth) def test_sandbox_client_rejects_ambiguous_bearer_configuration() -> None: auth = ClientCredentialsAuth( issuer="https://issuer.example.com", client_id="client", client_secret="secret", ) with pytest.raises(SandboxError, match="mutually exclusive"): SandboxClient( "localhost:50051", bearer_token="static", client_credentials=auth, ) def test_sandbox_client_close_releases_refresher_http_client( tmp_path: Path, monkeypatch: Any, ) -> None: """End-to-end check: an OIDC-backed client built by from_active_cluster() must close the refresher's httpx.Client when the SandboxClient is closed.""" gateway_name = "oidc-gw" gateway_dir = tmp_path / "openshell" / "gateways" / gateway_name mtls_dir = gateway_dir / "mtls" mtls_dir.mkdir(parents=True) (tmp_path / "openshell" / "active_gateway").write_text(gateway_name) (gateway_dir / "metadata.json").write_text( json.dumps( { "gateway_endpoint": "https://127.0.0.1:8443", "auth_mode": "oidc", } ) ) for f in ("ca.crt", "tls.crt", "tls.key"): (mtls_dir / f).write_text(f"-----BEGIN {f}-----\n-----END {f}-----\n") _write_bundle(gateway_dir, expires_at=int(time.time()) + 3600) monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) monkeypatch.delenv("OPENSHELL_GATEWAY", raising=False) # Capture the httpx.Client instance created inside the refresher by # monkey-patching _OidcRefresher to record it on construction. created: list[Any] = [] real_init = _OidcRefresher.__init__ def recording_init(self: Any, *args: Any, **kwargs: Any) -> None: real_init(self, *args, **kwargs) created.append(self._http) monkeypatch.setattr(_OidcRefresher, "__init__", recording_init) client = SandboxClient.from_active_cluster() assert len(created) == 1 http_client = created[0] assert not http_client.is_closed client.close() assert http_client.is_closed def test_refresher_concurrent_write_back_does_not_trample(tmp_path: Path) -> None: """Two writers calling `_write_to_disk` concurrently must each use their own tempfile (PID+random) and not corrupt each other's content. The final file must be valid JSON from exactly one of the writers, and no orphaned `.oidc_token..tmp` files should remain.""" _write_bundle(tmp_path, expires_at=int(time.time()) + 3600) r = _OidcRefresher(tmp_path, "g", write_back=False) try: N = 16 barrier = threading.Barrier(N) errors: list[BaseException] = [] def writer(idx: int) -> None: try: barrier.wait() r._write_to_disk( { "access_token": f"a-{idx}", "refresh_token": f"r-{idx}", "expires_at": 1_700_000_000 + idx, "issuer": DEFAULT_ISSUER, "client_id": "openshell-cli", } ) except BaseException as e: errors.append(e) threads = [threading.Thread(target=writer, args=(i,)) for i in range(N)] for t in threads: t.start() for t in threads: t.join() assert not errors, errors # Final file is valid JSON from one of the writers (race winner). final = json.loads((tmp_path / "oidc_token.json").read_text()) assert final["access_token"].startswith("a-") assert final["refresh_token"].startswith("r-") # No orphan tmp files left behind. mkstemp uses a random suffix # so each writer's tmp is distinct; the cleanup path on the # success branch is `.replace()`, which atomically moves the # tmp to the final path — no straggler tmp should remain. leftovers = sorted(tmp_path.glob(".oidc_token.*.tmp")) assert leftovers == [], f"orphan tmp files: {leftovers}" finally: r.close() class _WindowsPermissionError(PermissionError): winerror: int def test_atomic_replace_retries_windows_sharing_violations( tmp_path: Path, monkeypatch: Any ) -> None: source = tmp_path / "source" destination = tmp_path / "destination" source.write_text("new") destination.write_text("old") attempts = 0 delays: list[float] = [] real_replace = Path.replace def replace(path: Path, target: Path) -> Path: nonlocal attempts attempts += 1 if attempts < 3: error = _WindowsPermissionError("destination is busy") error.winerror = 32 raise error return real_replace(path, target) monkeypatch.setattr(sandbox_module, "_IS_WINDOWS", True) monkeypatch.setattr(Path, "replace", replace) monkeypatch.setattr(time, "sleep", delays.append) _atomic_replace(source, destination) assert attempts == 3 assert delays == [0.005, 0.01] assert destination.read_text() == "new" def test_atomic_replace_does_not_retry_permanent_windows_errors( tmp_path: Path, monkeypatch: Any ) -> None: source = tmp_path / "source" destination = tmp_path / "destination" source.write_text("new") attempts = 0 def replace(_path: Path, _target: Path) -> Path: nonlocal attempts attempts += 1 error = _WindowsPermissionError("access denied") error.winerror = 13 raise error monkeypatch.setattr(sandbox_module, "_IS_WINDOWS", True) monkeypatch.setattr(Path, "replace", replace) with pytest.raises(PermissionError, match="access denied"): _atomic_replace(source, destination) assert attempts == 1 def test_sandbox_wrapper_forwards_auth_kwargs_to_from_active_cluster( monkeypatch: Any, ) -> None: """The high-level `Sandbox` context manager must pass auto_refresh, write_back, and insecure through to SandboxClient.from_active_cluster so callers using the wrapper get parity with SandboxClient for OIDC-protected gateways.""" captured: dict[str, Any] = {} class _Sentinel(Exception): pass def fake_from_active_cluster(**kwargs: Any) -> Any: captured.update(kwargs) # Short-circuit the rest of __enter__ (which would try to create # a session against a real gateway). The kwargs we care about # have already been recorded. raise _Sentinel monkeypatch.setattr( SandboxClient, "from_active_cluster", staticmethod(fake_from_active_cluster) ) sandbox = Sandbox( workspace="default", cluster="my-gw", timeout=42.0, auto_refresh=False, write_back=False, insecure=True, ) import pytest as _pytest with _pytest.raises(_Sentinel): sandbox.__enter__() assert captured["cluster"] == "my-gw" assert captured["timeout"] == 42.0 assert captured["auto_refresh"] is False assert captured["write_back"] is False assert captured["insecure"] is True def test_sandbox_wrapper_defaults_match_from_active_cluster( monkeypatch: Any, ) -> None: """Sandbox(...) with no auth kwargs forwards the same defaults (auto_refresh=True, write_back=True, insecure=False) that SandboxClient.from_active_cluster uses, so the wrapper doesn't silently weaken the security posture.""" captured: dict[str, Any] = {} class _Sentinel(Exception): pass def fake_from_active_cluster(**kwargs: Any) -> Any: captured.update(kwargs) raise _Sentinel monkeypatch.setattr( SandboxClient, "from_active_cluster", staticmethod(fake_from_active_cluster) ) import pytest as _pytest with _pytest.raises(_Sentinel): Sandbox(workspace="default").__enter__() assert captured["auto_refresh"] is True assert captured["write_back"] is True assert captured["insecure"] is False # --------------------------------------------------------------------------- # Encoding regression tests (utf-8 explicit on all config file reads/writes) # --------------------------------------------------------------------------- def test_read_oidc_token_bundle_parses_non_ascii_utf8(tmp_path: Path) -> None: gateway_dir = tmp_path / "gw" gateway_dir.mkdir() payload = {"refresh_token": "tok", "issuer": "https://example.com/é"} (gateway_dir / "oidc_token.json").write_bytes( json.dumps(payload, ensure_ascii=False).encode("utf-8") ) result = _read_oidc_token_bundle(gateway_dir) assert result == payload def test_read_oidc_token_bundle_returns_none_on_corrupt_bytes(tmp_path: Path) -> None: gateway_dir = tmp_path / "gw" gateway_dir.mkdir() (gateway_dir / "oidc_token.json").write_bytes(b"\xff\xfe not utf-8") assert _read_oidc_token_bundle(gateway_dir) is None def test_load_cluster_bearer_token_handles_non_ascii_utf8_oidc(tmp_path: Path) -> None: gateway_dir = tmp_path / "gw" gateway_dir.mkdir() bundle = { "access_token": "accéss", "refresh_token": "ref", "expiry": "2099-01-01T00:00:00Z", "issuer": "https://example.com", "client_id": "c", "client_secret": "s", } (gateway_dir / "oidc_token.json").write_bytes( json.dumps(bundle, ensure_ascii=False).encode("utf-8") ) token = _load_cluster_bearer_token(gateway_dir) assert token == "accéss" def test_from_active_cluster_reads_utf8_bytes_from_active_gateway_and_metadata( tmp_path: Path, monkeypatch: Any, ) -> None: gateway_name = "gw-é" gateway_dir = tmp_path / "openshell" / "gateways" / gateway_name gateway_dir.mkdir(parents=True) (tmp_path / "openshell" / "active_gateway").write_bytes( gateway_name.encode("utf-8") ) meta = {"gateway_endpoint": "http://tést.example:8080"} (gateway_dir / "metadata.json").write_bytes( json.dumps(meta, ensure_ascii=False).encode("utf-8") ) monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) monkeypatch.delenv("OPENSHELL_GATEWAY", raising=False) client = SandboxClient.from_active_cluster() try: assert client._cluster_name == gateway_name assert client._endpoint == "tést.example:8080" finally: client.close() # ---- Sandbox label / selector API tests ---- def _make_sandbox_proto( id_: str, name: str, labels: dict[str, str] | None = None, phase: openshell_pb2.SandboxPhase = openshell_pb2.SANDBOX_PHASE_READY, version: int = 0, workspace: str = "default", ) -> openshell_pb2.Sandbox: sandbox = openshell_pb2.Sandbox() sandbox.metadata.id = id_ sandbox.metadata.name = name sandbox.metadata.workspace = workspace for key, value in (labels or {}).items(): sandbox.metadata.labels[key] = value sandbox.status.phase = phase sandbox.status.current_policy_version = version return sandbox def _make_workload_template_proto( name: str, *, workspace: str = "default", ) -> openshell_pb2.SandboxWorkloadTemplate: template = openshell_pb2.SandboxWorkloadTemplate() template.metadata.name = name template.metadata.workspace = workspace template.spec.workload.image = f"ghcr.io/test/{name}:latest" template.spec.workload.resources.cpu = "1" template.spec.workload.resources.memory = "512Mi" return template class _FakeSandboxStub: def __init__( self, listed: list[openshell_pb2.Sandbox] | None = None, listed_pages: list[list[openshell_pb2.Sandbox]] | None = None, ) -> None: self.create_request: openshell_pb2.CreateSandboxRequest | None = None self.list_request: openshell_pb2.ListSandboxesRequest | None = None self.get_request: openshell_pb2.GetSandboxRequest | None = None self.delete_request: openshell_pb2.DeleteSandboxRequest | None = None self.stop_request: openshell_pb2.StopSandboxRequest | None = None self.start_request: openshell_pb2.StartSandboxRequest | None = None self.create_template_request: ( openshell_pb2.CreateSandboxTemplateRequest | None ) = None self.get_template_request: openshell_pb2.GetSandboxTemplateRequest | None = None self.list_template_request: openshell_pb2.ListSandboxTemplatesRequest | None = ( None ) self.delete_template_request: ( openshell_pb2.DeleteSandboxTemplateRequest | None ) = None self._listed = listed or [] self._listed_pages = listed_pages self.list_requests: list[openshell_pb2.ListSandboxesRequest] = [] self._templates: list[openshell_pb2.SandboxWorkloadTemplate] = [] def GetSandbox( self, request: openshell_pb2.GetSandboxRequest, timeout: float | None = None, ) -> Any: self.get_request = request _ = timeout return SimpleNamespace( sandbox=_make_sandbox_proto( "sandbox-1", _request_sandbox(request), workspace=_request_workspace(request) or "default", ) ) def DeleteSandbox( self, request: openshell_pb2.DeleteSandboxRequest, timeout: float | None = None, ) -> Any: self.delete_request = request _ = timeout return SimpleNamespace(outcome=1, sandbox_id="sb-1") def StopSandbox( self, request: openshell_pb2.StopSandboxRequest, timeout: float | None = None, ) -> Any: self.stop_request = request _ = timeout return SimpleNamespace( sandbox=_make_sandbox_proto( "sandbox-1", _request_sandbox(request), phase=openshell_pb2.SANDBOX_PHASE_STOPPED, workspace=_request_workspace(request) or "default", ) ) def StartSandbox( self, request: openshell_pb2.StartSandboxRequest, timeout: float | None = None, ) -> Any: self.start_request = request _ = timeout return SimpleNamespace( sandbox=_make_sandbox_proto( "sandbox-1", _request_sandbox(request), phase=openshell_pb2.SANDBOX_PHASE_STARTING, workspace=_request_workspace(request) or "default", ) ) def CreateSandbox( self, request: openshell_pb2.CreateSandboxRequest, timeout: float | None = None, ) -> Any: self.create_request = request _ = timeout return SimpleNamespace( sandbox=_make_sandbox_proto( "sandbox-1", request.name or "generated", dict(request.labels), workspace=_request_workspace(request) or "default", ), service_urls={ exposure.service: f"https://{exposure.service}.example.test/" for exposure in request.service_exposures }, ) def ListSandboxes( self, request: openshell_pb2.ListSandboxesRequest, timeout: float | None = None, ) -> Any: self.list_request = request self.list_requests.append(deepcopy(request)) _ = timeout if self._listed_pages is not None: page = int(request.page_token or "0") next_page_token = ( str(page + 1) if page + 1 < len(self._listed_pages) else "" ) return SimpleNamespace( sandboxes=list(self._listed_pages[page]), next_page_token=next_page_token, ) return SimpleNamespace(sandboxes=list(self._listed)) def CreateSandboxTemplate( self, request: openshell_pb2.CreateSandboxTemplateRequest, timeout: float | None = None, ) -> Any: self.create_template_request = request _ = timeout self._templates.append(request.template) return SimpleNamespace(template=request.template) def GetSandboxTemplate( self, request: openshell_pb2.GetSandboxTemplateRequest, timeout: float | None = None, ) -> Any: self.get_template_request = request _ = timeout return SimpleNamespace( template=_make_workload_template_proto( request.name, workspace=_request_workspace(request) or "default", ) ) def ListSandboxTemplates( self, request: openshell_pb2.ListSandboxTemplatesRequest, timeout: float | None = None, ) -> Any: self.list_template_request = request _ = timeout return SimpleNamespace(templates=list(self._templates)) def DeleteSandboxTemplate( self, request: openshell_pb2.DeleteSandboxTemplateRequest, timeout: float | None = None, ) -> Any: self.delete_template_request = request _ = timeout return SimpleNamespace(outcome=1, sandbox_id="sb-1") class _RecordingHighLevelClient: """A stand-in for SandboxClient used to observe high-level forwarding.""" def __init__(self) -> None: self.create_kwargs: dict[str, Any] | None = None self.create_template_kwargs: dict[str, Any] | None = None def create_session( self, *, workspace: str, spec: Any = None, name: str | None = None, labels: Any = None, ) -> Any: self.create_kwargs = { "workspace": workspace, "spec": spec, "name": name, "labels": labels, } return SimpleNamespace(sandbox=SimpleNamespace(name=name or "generated")) def create_session_from_template( self, *, workspace: str, workload_template: str, spec: Any = None, name: str | None = None, labels: Any = None, ) -> Any: self.create_template_kwargs = { "workspace": workspace, "workload_template": workload_template, "spec": spec, "name": name, "labels": labels, } return SimpleNamespace(sandbox=SimpleNamespace(name=name or "generated")) def wait_ready( self, name: str, *, workspace: str, timeout_seconds: float = 300.0 ) -> SandboxRef: _ = timeout_seconds return SandboxRef( id="sandbox-1", name=name, workspace=workspace, status=SandboxStatusRef(phase=2, current_policy_version=0), ) def test_create_forwards_name_and_labels() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) ref = client.create( workspace="default", name="job-1", labels={"aiq": "deep-research"} ) assert stub.create_request is not None assert stub.create_request.name == "job-1" assert dict(stub.create_request.labels) == {"aiq": "deep-research"} assert dict(ref.labels) == {"aiq": "deep-research"} def test_create_forwards_service_exposures() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) ref = client.create( workspace="default", name="app-server", service_exposures=[ ServiceExposure(target_port=4500), ServiceExposure( service="metrics", target_port=9090, authorization_mode=ServiceAuthorizationMode.BEARER_PASSTHROUGH, ), ], ) assert stub.create_request is not None assert [ (exposure.service, exposure.target_port, exposure.authorization_mode) for exposure in stub.create_request.service_exposures ] == [ ("", 4500, openshell_pb2.SERVICE_AUTHORIZATION_MODE_STRIP), ( "metrics", 9090, openshell_pb2.SERVICE_AUTHORIZATION_MODE_BEARER_PASSTHROUGH, ), ] assert dict(ref.service_urls) == { "": "https://.example.test/", "metrics": "https://metrics.example.test/", } def test_create_from_template_forwards_workload_template() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) spec = openshell_pb2.SandboxSpec( providers=["github"], command=["/opt/worker", "--serve"], tty=True, ) ref = client.create_from_template( workspace="default", workload_template="gpu-kata", spec=spec, name="job-1", labels={"team": "runtime"}, ) assert stub.create_request is not None assert stub.create_request.name == "job-1" assert stub.create_request.workload_template == "gpu-kata" assert dict(stub.create_request.labels) == {"team": "runtime"} assert list(stub.create_request.spec.providers) == ["github"] assert list(stub.create_request.spec.command) == ["/opt/worker", "--serve"] assert stub.create_request.spec.tty is True assert dict(ref.labels) == {"team": "runtime"} def test_create_from_template_rejects_empty_workload_template() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) with pytest.raises(SandboxError): client.create_from_template(workspace="default", workload_template=" ") assert stub.create_request is None def test_sandbox_template_create_builds_template_from_public_fields() -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) created = client.create( workspace="default", name="gpu-kata", image="ghcr.io/test/gpu-kata:latest", labels={"team": "runtime"}, annotations={"owner": "platform"}, environment={"FEATURE_FLAG": "on"}, cpu="1", memory="512Mi", gpu_count=2, driver_config={"kubernetes": {"runtime_class_name": "kata"}}, ) assert created.metadata.name == "gpu-kata" assert stub.create_template_request is not None assert _request_workspace(stub.create_template_request) == "default" template = stub.create_template_request.template assert template.metadata.name == "gpu-kata" assert dict(template.metadata.labels) == {"team": "runtime"} assert dict(template.metadata.annotations) == {"owner": "platform"} assert template.spec.workload.image == "ghcr.io/test/gpu-kata:latest" assert dict(template.spec.workload.environment) == {"FEATURE_FLAG": "on"} assert template.spec.workload.resources.cpu == "1" assert template.spec.workload.resources.memory == "512Mi" assert template.spec.workload.resources.gpu.count == 2 assert template.spec.driver_config["kubernetes"]["runtime_class_name"] == "kata" def test_sandbox_template_create_materializes_default_workload() -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) client.create(workspace="default", name="base") assert stub.create_template_request is not None template = stub.create_template_request.template assert template.HasField("spec") assert template.spec.HasField("workload") assert template.spec.workload.image == "" def test_sandbox_template_create_materializes_workload_with_driver_config_only() -> ( None ): stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) client.create( workspace="default", name="kata-default-image", driver_config={"kubernetes": {"runtime_class_name": "kata"}}, ) assert stub.create_template_request is not None template = stub.create_template_request.template assert template.HasField("spec") assert template.spec.HasField("workload") assert template.spec.workload.image == "" assert template.spec.driver_config["kubernetes"]["runtime_class_name"] == "kata" def test_sandbox_template_create_rejects_missing_public_name() -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) with pytest.raises(SandboxError): client.create(workspace="default", image="ghcr.io/test/python:latest") assert stub.create_template_request is None def test_sandbox_template_create_rejects_template_and_builder_fields() -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) template = _make_workload_template_proto("gpu-kata") with pytest.raises(SandboxError): client.create(workspace="default", template=template, image="override") assert stub.create_template_request is None @pytest.mark.parametrize( "builder_kwargs", ( {"labels": {}}, {"annotations": {}}, {"environment": {}}, {"driver_config": {}}, ), ) def test_sandbox_template_create_allows_template_and_empty_builder_mappings( builder_kwargs: dict[str, Any], ) -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) template = _make_workload_template_proto("gpu-kata") created = client.create( workspace="default", template=template, **builder_kwargs, ) assert created.metadata.name == "gpu-kata" assert stub.create_template_request is not None assert stub.create_template_request.template.metadata.name == template.metadata.name assert ( stub.create_template_request.template.spec.workload.image == template.spec.workload.image ) @pytest.mark.parametrize( "builder_kwargs", ( {"labels": {"team": "runtime"}}, {"annotations": {"owner": "platform"}}, {"environment": {"FEATURE_FLAG": "on"}}, {"driver_config": {"kubernetes": {"runtime_class_name": "kata"}}}, ), ) def test_sandbox_template_create_rejects_template_and_non_empty_builder_mappings( builder_kwargs: dict[str, Any], ) -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) template = _make_workload_template_proto("gpu-kata") with pytest.raises(SandboxError): client.create( workspace="default", template=template, **builder_kwargs, ) assert stub.create_template_request is None def test_sandbox_template_create_rejects_non_positive_gpu_count() -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) with pytest.raises(SandboxError): client.create(workspace="default", name="gpu-kata", gpu_count=0) assert stub.create_template_request is None def test_sandbox_template_client_crud_forwards_requests() -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) template = _make_workload_template_proto("gpu-kata") template.spec.driver_config.update({"kubernetes": {"runtime_class_name": "kata"}}) created = client.create(workspace="default", template=template) assert created.metadata.name == "gpu-kata" assert stub.create_template_request is not None assert _request_workspace(stub.create_template_request) == "default" assert ( stub.create_template_request.template.spec.workload.image == "ghcr.io/test/gpu-kata:latest" ) assert ( stub.create_template_request.template.spec.driver_config["kubernetes"][ "runtime_class_name" ] == "kata" ) got = client.get("gpu-kata", workspace="default") assert got.metadata.name == "gpu-kata" assert stub.get_template_request is not None assert stub.get_template_request.name == "gpu-kata" assert _request_workspace(stub.get_template_request) == "default" listed = client.list_all( workspace="default", page_size=50, label_selector="team=runtime" ) assert len(listed) == 1 assert stub.list_template_request is not None assert _request_workspace(stub.list_template_request) == "default" assert stub.list_template_request.page_size == 50 assert stub.list_template_request.page_token == "" assert stub.list_template_request.label_selector == "team=runtime" assert not _request_selects_all_workspaces(stub.list_template_request) assert client.delete("gpu-kata", workspace="default").outcome == 1 assert stub.delete_template_request is not None assert stub.delete_template_request.name == "gpu-kata" assert _request_workspace(stub.delete_template_request) == "default" def test_sandbox_template_list_for_all_workspaces_selects_all() -> None: stub = _FakeSandboxStub() client = _template_client_with_fake_stub(stub) client.list_all_for_all_workspaces(page_size=100, label_selector="team=runtime") assert stub.list_template_request is not None assert _request_selects_all_workspaces(stub.list_template_request) assert _request_workspace(stub.list_template_request) is None assert stub.list_template_request.page_size == 100 assert stub.list_template_request.page_token == "" assert stub.list_template_request.label_selector == "team=runtime" def test_stop_and_start_forward_workspace_and_return_phase() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) stopped = client.stop("job-1", workspace="team-a") assert stub.stop_request is not None assert _request_sandbox(stub.stop_request) == "job-1" assert _request_workspace(stub.stop_request) == "team-a" assert stopped.phase == openshell_pb2.SANDBOX_PHASE_STOPPED starting = client.start("job-1", workspace="team-a") assert stub.start_request is not None assert _request_sandbox(stub.start_request) == "job-1" assert _request_workspace(stub.start_request) == "team-a" assert starting.phase == openshell_pb2.SANDBOX_PHASE_STARTING @pytest.mark.parametrize( ("phase", "should_succeed"), [ (openshell_pb2.SANDBOX_PHASE_COMPLETED, True), (openshell_pb2.SANDBOX_PHASE_ERROR, False), ], ) def test_wait_ready_handles_terminal_main_process_results( phase: openshell_pb2.SandboxPhase, should_succeed: bool ) -> None: class TerminalStub(_FakeSandboxStub): def GetSandbox( self, request: openshell_pb2.GetSandboxRequest, timeout: float | None = None, ) -> Any: _ = timeout return SimpleNamespace( sandbox=_make_sandbox_proto( "sandbox-1", _request_sandbox(request), phase=phase, workspace=_request_workspace(request) or "default", ) ) client = _client_with_fake_stub(TerminalStub()) if should_succeed: result = client.wait_ready("job-1", workspace="default", timeout_seconds=0.1) assert result.phase == openshell_pb2.SANDBOX_PHASE_COMPLETED else: with pytest.raises(SandboxError, match="entered error phase"): client.wait_ready("job-1", workspace="default", timeout_seconds=0.1) def test_create_without_args_sends_empty_metadata() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) client.create(workspace="default") assert stub.create_request is not None assert stub.create_request.name == "" assert dict(stub.create_request.labels) == {} assert _request_workspace(stub.create_request) == "default" def test_create_copies_caller_labels() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) caller_labels = {"aiq": "deep-research"} client.create(workspace="default", labels=caller_labels) caller_labels["aiq"] = "mutated" assert stub.create_request is not None assert dict(stub.create_request.labels) == {"aiq": "deep-research"} def test_create_session_forwards_name_and_labels() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) session = client.create_session( workspace="default", name="job-2", labels={"team": "aiq"} ) assert stub.create_request is not None assert stub.create_request.name == "job-2" assert dict(stub.create_request.labels) == {"team": "aiq"} assert session.sandbox.name == "job-2" def test_list_forwards_label_selector() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) client.list_all(workspace="default", label_selector="aiq=deep-research") assert stub.list_request is not None assert stub.list_request.label_selector == "aiq=deep-research" assert _request_workspace(stub.list_request) == "default" def test_list_without_selector_sends_empty_string() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) client.list_all(workspace="default") assert stub.list_request is not None assert stub.list_request.label_selector == "" def test_list_follows_continuation_tokens() -> None: stub = _FakeSandboxStub( listed_pages=[ [_make_sandbox_proto("sandbox-1", "job-1")], [_make_sandbox_proto("sandbox-2", "job-2")], ] ) client = _client_with_fake_stub(stub) pager = client.list(workspace="default", page_size=1, label_selector="team=core") assert stub.list_requests == [] first = next(pager) assert [sandbox.name for sandbox in first.items] == ["job-1"] assert first.next_page_token == "1" second = next(pager) assert [sandbox.name for sandbox in second.items] == ["job-2"] assert second.next_page_token == "" with pytest.raises(StopIteration): next(pager) assert len(stub.list_requests) == 2 assert stub.list_requests[0].page_token == "" assert stub.list_requests[1].page_token == "1" assert stub.list_requests[1].label_selector == "team=core" def test_list_passes_initial_page_token() -> None: stub = _FakeSandboxStub( listed_pages=[ [_make_sandbox_proto("sandbox-1", "skipped")], [_make_sandbox_proto("sandbox-2", "resumed")], ] ) client = _client_with_fake_stub(stub) page = next(client.list(workspace="default", page_token="1")) assert [sandbox.name for sandbox in page.items] == ["resumed"] assert stub.list_requests[0].page_token == "1" def test_pager_retries_same_token_after_fetch_error() -> None: tokens: list[str] = [] def fetch(token: str) -> Page[int]: tokens.append(token) if len(tokens) == 1: raise RuntimeError("temporary failure") return Page(items=[1], next_page_token="") pager = Pager(fetch, page_token="resume") with pytest.raises(RuntimeError, match="temporary failure"): next(pager) assert next(pager).items == [1] assert tokens == ["resume", "resume"] def test_pager_rejects_a_repeated_continuation_token() -> None: pager = Pager( lambda token: Page(items=[1], next_page_token=token), page_token="resume", ) with pytest.raises(SandboxError, match="repeated continuation token"): next(pager) def test_pager_bounds_consumed_token_count(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(sandbox_module, "_PAGER_MAX_CONSUMED_TOKENS", 1) requests: list[str] = [] pager = Pager( lambda token: ( requests.append(token), Page(items=[token], next_page_token="next"), )[1], page_token="first", ) assert next(pager).items == ["first"] with pytest.raises(SandboxError, match="token history limit exceeded"): next(pager) assert requests == ["first"] def test_pager_bounds_consumed_token_bytes(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(sandbox_module, "_PAGER_MAX_CONSUMED_TOKEN_BYTES", 1) requests: list[str] = [] pager = Pager( lambda token: ( requests.append(token), Page(items=[token], next_page_token="next"), )[1], page_token="too-large", ) with pytest.raises(SandboxError, match="token history limit exceeded"): next(pager) assert requests == [] def test_list_ids_forwards_label_selector() -> None: stub = _FakeSandboxStub(listed=[_make_sandbox_proto("sandbox-1", "job-1")]) client = _client_with_fake_stub(stub) ids = client.list_ids(workspace="default", label_selector="aiq=deep-research") assert stub.list_request is not None assert stub.list_request.label_selector == "aiq=deep-research" assert ids == ["sandbox-1"] def test_sandbox_ref_retains_gateway_labels() -> None: proto = _make_sandbox_proto( "sandbox-1", "job-1", {"aiq": "deep-research", "env": "dev"} ) ref = _sandbox_ref(proto) assert dict(ref.labels) == {"aiq": "deep-research", "env": "dev"} def test_sandbox_ref_includes_main_process_result() -> None: proto = _make_sandbox_proto("sandbox-1", "job-1") proto.status.exit_code = 0 proto.status.restart_count = 2 proto.status.next_restart_time.FromMilliseconds(1_700_000_000_000) proto.status.main_process_started_time.FromMilliseconds(1_699_999_000_000) status = _sandbox_ref(proto).status assert status.exit_code == 0 assert status.restart_count == 2 assert status.next_restart_at_ms == 1_700_000_000_000 assert status.main_process_started_at_ms == 1_699_999_000_000 def test_returned_labels_are_immutable() -> None: proto = _make_sandbox_proto("sandbox-1", "job-1", {"aiq": "deep-research"}) ref = _sandbox_ref(proto) with pytest.raises(TypeError): ref.labels["mutated"] = "nope" # type: ignore[index] def test_direct_sandbox_ref_construction_defaults_labels() -> None: ref = SandboxRef( id="sandbox-1", name="job-1", workspace="default", status=SandboxStatusRef(phase=2, current_policy_version=0), ) assert dict(ref.labels) == {} def test_sandbox_ref_stays_hashable_with_labels_excluded_from_identity() -> None: ref_a = _sandbox_ref(_make_sandbox_proto("sandbox-1", "job-1", {"aiq": "a"})) ref_b = _sandbox_ref(_make_sandbox_proto("sandbox-1", "job-1", {"aiq": "b"})) # Frozen dataclass must remain hashable despite the immutable labels field. assert hash(ref_a) == hash(ref_b) # Labels are excluded from identity: same (id, name, status) compares equal. assert ref_a == ref_b assert {ref_a, ref_b} == {ref_a} def test_sandbox_ref_labels_support_standard_serialization() -> None: ref = _sandbox_ref( _make_sandbox_proto("sandbox-1", "job-1", {"aiq": "deep-research"}) ) assert asdict(ref)["labels"] == {"aiq": "deep-research"} assert dict(deepcopy(ref).labels) == {"aiq": "deep-research"} assert dict(pickle.loads(pickle.dumps(ref)).labels) == {"aiq": "deep-research"} def test_default_sandbox_ref_labels_support_standard_serialization() -> None: ref = SandboxRef( id="sandbox-1", name="job-1", workspace="default", status=SandboxStatusRef(phase=2, current_policy_version=0), ) assert asdict(ref)["labels"] == {} assert dict(deepcopy(ref).labels) == {} assert dict(pickle.loads(pickle.dumps(ref)).labels) == {} def test_direct_sandbox_ref_copies_and_freezes_labels() -> None: labels = {"aiq": "deep-research"} ref = SandboxRef( id="sandbox-1", name="job-1", workspace="default", status=SandboxStatusRef(phase=2, current_policy_version=0), labels=labels, ) labels["aiq"] = "mutated" assert dict(ref.labels) == {"aiq": "deep-research"} with pytest.raises(TypeError): ref.labels["mutated"] = "nope" # type: ignore[index] def test_high_level_creation_forwards_name_and_labels( monkeypatch: pytest.MonkeyPatch, ) -> None: recording = _RecordingHighLevelClient() monkeypatch.setattr( SandboxClient, "from_active_cluster", classmethod(lambda _cls, **_kwargs: recording), ) sandbox = Sandbox( workspace="staging", name="job-1", labels={"aiq": "deep-research"}, delete_on_exit=False, ) sandbox.__enter__() assert recording.create_kwargs == { "workspace": "staging", "spec": None, "name": "job-1", "labels": {"aiq": "deep-research"}, } def test_high_level_template_creation_forwards_workload_template( monkeypatch: pytest.MonkeyPatch, ) -> None: recording = _RecordingHighLevelClient() monkeypatch.setattr( SandboxClient, "from_active_cluster", classmethod(lambda _cls, **_kwargs: recording), ) spec = openshell_pb2.SandboxSpec( providers=["github"], command=["/opt/worker", "--serve"], tty=True, ) sandbox = Sandbox( workspace="staging", workload_template="gpu-kata", spec=spec, name="job-1", labels={"team": "runtime"}, delete_on_exit=False, ) sandbox.__enter__() assert recording.create_template_kwargs == { "workspace": "staging", "workload_template": "gpu-kata", "spec": spec, "name": "job-1", "labels": {"team": "runtime"}, } assert recording.create_template_kwargs is not None forwarded_spec = recording.create_template_kwargs["spec"] assert list(forwarded_spec.command) == ["/opt/worker", "--serve"] assert forwarded_spec.tty is True def test_high_level_attach_rejects_name() -> None: sandbox = Sandbox(workspace="default", sandbox="existing-sandbox", name="job-1") with pytest.raises(SandboxError): sandbox.__enter__() def test_high_level_attach_rejects_labels() -> None: ref = SandboxRef( id="sandbox-1", name="existing", workspace="default", status=SandboxStatusRef(phase=2, current_policy_version=0), ) sandbox = Sandbox(workspace="default", sandbox=ref, labels={"aiq": "deep-research"}) with pytest.raises(SandboxError): sandbox.__enter__() def test_high_level_attach_rejects_workload_template() -> None: sandbox = Sandbox( workspace="default", sandbox="existing-sandbox", workload_template="gpu-kata" ) with pytest.raises(SandboxError): sandbox.__enter__() # --------------------------------------------------------------------------- # Workspace support # --------------------------------------------------------------------------- def test_create_passes_workspace_to_proto() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) ref = client.create(workspace="staging", name="job-1") assert stub.create_request is not None assert _request_workspace(stub.create_request) == "staging" assert ref.workspace == "staging" def test_get_passes_workspace_to_proto() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) ref = client.get("job-1", workspace="production") assert stub.get_request is not None assert _request_workspace(stub.get_request) == "production" assert ref.workspace == "production" def test_delete_passes_workspace_to_proto() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) result = client.delete("job-1", workspace="staging") assert result.outcome == 1 assert stub.delete_request is not None assert _request_workspace(stub.delete_request) == "staging" assert not stub.delete_request.allow_missing @pytest.mark.parametrize("outcome", [0, 1, 2, 3, 99]) def test_delete_preserves_outcome_and_identity(outcome: int) -> None: class Stub: def DeleteSandbox(self, request: Any, **_kwargs: Any) -> Any: assert request.allow_missing return openshell_pb2.DeleteSandboxResponse( outcome=cast("openshell_pb2.DeletionOutcome", outcome), sandbox_id="original-id", ) result = _client_with_fake_stub(Stub()).delete( "job", workspace="default", allow_missing=True ) assert int(result.outcome) == outcome assert result.sandbox_id == "original-id" if outcome == 99: assert result.outcome not in ( DeletionOutcome.COMPLETED, DeletionOutcome.ALREADY_ABSENT, ) def test_list_for_all_workspaces_sets_flag() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) client.list_all_for_all_workspaces() assert stub.list_request is not None assert _request_selects_all_workspaces(stub.list_request) assert _request_workspace(stub.list_request) is None def test_list_with_workspace_passes_workspace() -> None: stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) client.list_all(workspace="staging") assert stub.list_request is not None assert _request_workspace(stub.list_request) == "staging" assert not _request_selects_all_workspaces(stub.list_request) def test_sandbox_ref_includes_workspace_from_proto() -> None: proto = _make_sandbox_proto("sandbox-1", "job-1", workspace="production") ref = _sandbox_ref(proto) assert ref.workspace == "production" def test_sandbox_ref_includes_workload_template_provenance() -> None: proto = _make_sandbox_proto("sandbox-1", "job-1") proto.created_from_workload_template.name = "gpu-kata" proto.created_from_workload_template.resource_version = "7" ref = _sandbox_ref(proto) assert ref.created_from_workload_template is not None assert ref.created_from_workload_template.name == "gpu-kata" assert ref.created_from_workload_template.resource_version == "7" def test_sandbox_session_delete_passes_workspace() -> None: from openshell.sandbox import SandboxSession stub = _FakeSandboxStub() client = _client_with_fake_stub(stub) ref = SandboxRef( id="sandbox-1", name="job-1", workspace="staging", status=SandboxStatusRef(phase=2, current_policy_version=0), ) session = SandboxSession(client, ref) session.delete() assert stub.delete_request is not None assert _request_workspace(stub.delete_request) == "staging"