mirror of
https://github.com/NVIDIA/OpenShell.git
synced 2026-10-02 07:34:45 +08:00
feat(policy): add agentic approval loop (#1528)
This commit is contained in:
+35
-12
@@ -12,8 +12,9 @@ import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from collections import namedtuple
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import grpc
|
||||
@@ -26,6 +27,16 @@ from ._proto import (
|
||||
openshell_pb2_grpc,
|
||||
)
|
||||
|
||||
_ClientCallDetailsBase = namedtuple(
|
||||
"_ClientCallDetailsBase",
|
||||
("method", "timeout", "metadata", "credentials", "wait_for_ready", "compression"),
|
||||
)
|
||||
|
||||
|
||||
class _ClientCallDetails(_ClientCallDetailsBase, grpc.ClientCallDetails):
|
||||
pass
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import builtins
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
@@ -73,21 +84,33 @@ class _BearerAuthInterceptor(
|
||||
self._token_provider = token_provider
|
||||
|
||||
def _attach(self, details: grpc.ClientCallDetails) -> grpc.ClientCallDetails:
|
||||
metadata = list(details.metadata) if details.metadata else []
|
||||
original_metadata = getattr(details, "metadata", None)
|
||||
metadata = list(original_metadata) if original_metadata else []
|
||||
metadata.append(("authorization", f"Bearer {self._token_provider()}"))
|
||||
return details._replace(metadata=metadata)
|
||||
return _ClientCallDetails(
|
||||
getattr(details, "method", None),
|
||||
getattr(details, "timeout", None),
|
||||
metadata,
|
||||
getattr(details, "credentials", None),
|
||||
getattr(details, "wait_for_ready", None),
|
||||
getattr(details, "compression", None),
|
||||
)
|
||||
|
||||
def intercept_unary_unary(self, continuation, details, request):
|
||||
return continuation(self._attach(details), request)
|
||||
def intercept_unary_unary(self, continuation, client_call_details, request):
|
||||
return continuation(self._attach(client_call_details), request)
|
||||
|
||||
def intercept_unary_stream(self, continuation, details, request):
|
||||
return continuation(self._attach(details), request)
|
||||
def intercept_unary_stream(self, continuation, client_call_details, request):
|
||||
return continuation(self._attach(client_call_details), request)
|
||||
|
||||
def intercept_stream_unary(self, continuation, details, request_iterator):
|
||||
return continuation(self._attach(details), request_iterator)
|
||||
def intercept_stream_unary(
|
||||
self, continuation, client_call_details, request_iterator
|
||||
):
|
||||
return continuation(self._attach(client_call_details), request_iterator)
|
||||
|
||||
def intercept_stream_stream(self, continuation, details, request_iterator):
|
||||
return continuation(self._attach(details), request_iterator)
|
||||
def intercept_stream_stream(
|
||||
self, continuation, client_call_details, request_iterator
|
||||
):
|
||||
return continuation(self._attach(client_call_details), request_iterator)
|
||||
|
||||
|
||||
def _normalize_bearer(
|
||||
@@ -96,7 +119,7 @@ def _normalize_bearer(
|
||||
if bearer is None:
|
||||
return None
|
||||
if callable(bearer):
|
||||
return bearer
|
||||
return cast("Callable[[], str]", bearer)
|
||||
token = bearer
|
||||
return lambda: token
|
||||
|
||||
|
||||
@@ -174,15 +174,13 @@ class _FakeClientCallDetails:
|
||||
self.wait_for_ready = wait_for_ready
|
||||
|
||||
def _replace(self, **kwargs: Any) -> _FakeClientCallDetails:
|
||||
current = {
|
||||
"method": self.method,
|
||||
"timeout": self.timeout,
|
||||
"metadata": self.metadata,
|
||||
"credentials": self.credentials,
|
||||
"wait_for_ready": self.wait_for_ready,
|
||||
}
|
||||
current.update(kwargs)
|
||||
return _FakeClientCallDetails(**current)
|
||||
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:
|
||||
@@ -947,7 +945,7 @@ def test_refresher_concurrent_calls_share_one_refresh(tmp_path: Path) -> None:
|
||||
)
|
||||
|
||||
r = _OidcRefresher(tmp_path, "g", write_back=False)
|
||||
_install_mock_transport(r, _httpx.MockTransport(handler)) # type: ignore[has-type]
|
||||
_install_mock_transport(r, _httpx.MockTransport(handler))
|
||||
|
||||
results: list[str] = []
|
||||
errors: list[BaseException] = []
|
||||
|
||||
Reference in New Issue
Block a user