mirror of
https://github.com/NVIDIA/OpenShell.git
synced 2026-10-02 15:40:03 +08:00
* feat(service): add bearer authorization passthrough Signed-off-by: Derek Carr <decarr@redhat.com> * docs(sdk): add service authorization migration guide Signed-off-by: Derek Carr <decarr@redhat.com> * fix(server): remove stale version import Signed-off-by: Derek Carr <decarr@redhat.com> * docs(upgrade): remove service authorization SDK guide Signed-off-by: Derek Carr <decarr@redhat.com> * fix(e2e): relabel provider readiness TLS mount Signed-off-by: Derek Carr <decarr@redhat.com> * test(e2e): stabilize exposed service routing Signed-off-by: Derek Carr <decarr@redhat.com> * test(e2e): support HTTPS service routing Signed-off-by: Derek Carr <decarr@redhat.com> --------- Signed-off-by: Derek Carr <decarr@redhat.com>
3028 lines
100 KiB
Python
3028 lines
100 KiB
Python
# 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.<rand>.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"
|