mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
test(anthropic): cover the token exchange and Keycloak default posters
Both posters built their HTTPHandler inline, so nothing exercised the guarantees that matter there: redirects stay disabled, an HTTPStatusError comes back as its response rather than escaping, and a None response raises instead of being dereferenced. Each now takes a handler factory that defaults to the real one, which is the same injection seam the rest of this package already uses for its poster and secret reader Coverage for litellm/llms/base_llm/auth goes from 90% to 97%: client_credentials from 79% to 98%, with only its assert_never arm left, and token_exchange from 90% to 96%
This commit is contained in:
parent
cc18f804f8
commit
498592e0ce
4 changed files with 205 additions and 12 deletions
|
|
@ -49,24 +49,29 @@ def _default_secret_reader(ref: str) -> str | None:
|
|||
return get_secret_str(ref)
|
||||
|
||||
|
||||
def _new_keycloak_handler() -> "HTTPHandler":
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
handler: Final = HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0))
|
||||
handler.client.follow_redirects = False
|
||||
return handler
|
||||
|
||||
|
||||
class _HttpxSyncKeycloakPoster:
|
||||
"""Dedicated HTTPHandler for the Keycloak token POST: no ``logging_obj`` (so litellm's
|
||||
request/response logging never sees the client secret or the fetched token), redirects
|
||||
disabled. A separate instance from the outer engine's own poster, since this is a genuinely
|
||||
new HTTP call site whose no-logging guarantee must be built here, not assumed inherited."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, handler_factory: Callable[[], "HTTPHandler"] = _new_keycloak_handler) -> None:
|
||||
self._lock: Final = threading.Lock()
|
||||
self._handler_factory: Final = handler_factory
|
||||
self._handler: HTTPHandler | None = None
|
||||
|
||||
def _handler_instance(self) -> "HTTPHandler":
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
with self._lock:
|
||||
if self._handler is None:
|
||||
handler: Final = HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0))
|
||||
handler.client.follow_redirects = False
|
||||
self._handler = handler
|
||||
self._handler = self._handler_factory()
|
||||
return self._handler
|
||||
|
||||
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
|
||||
|
|
|
|||
|
|
@ -233,23 +233,28 @@ def _default_assertion_reader(ref: str) -> str | None:
|
|||
return get_secret_str(ref)
|
||||
|
||||
|
||||
def _new_exchange_handler() -> "HTTPHandler":
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
handler: Final = HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0))
|
||||
handler.client.follow_redirects = False
|
||||
return handler
|
||||
|
||||
|
||||
class _HttpxSyncTokenPoster:
|
||||
"""Default poster: a dedicated HTTPHandler (no logging_obj, so litellm's
|
||||
pre/post-call body logging never sees the exchange POST); returns the
|
||||
response for any status."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, handler_factory: Callable[[], "HTTPHandler"] = _new_exchange_handler) -> None:
|
||||
self._lock: Final = threading.Lock()
|
||||
self._handler_factory: Final = handler_factory
|
||||
self._handler: HTTPHandler | None = None
|
||||
|
||||
def _handler_instance(self) -> "HTTPHandler":
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
with self._lock:
|
||||
if self._handler is None:
|
||||
handler: Final = HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0))
|
||||
handler.client.follow_redirects = False
|
||||
self._handler = handler
|
||||
self._handler = self._handler_factory()
|
||||
return self._handler
|
||||
|
||||
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
|
||||
|
|
|
|||
|
|
@ -8,10 +8,14 @@ import httpx
|
|||
import pytest
|
||||
|
||||
from litellm.llms.base_llm.auth.client_credentials import (
|
||||
_HttpxSyncKeycloakPoster,
|
||||
_default_secret_reader,
|
||||
_new_keycloak_handler,
|
||||
fetch_keycloak_assertion,
|
||||
keycloak_assertion_source,
|
||||
)
|
||||
from litellm.llms.base_llm.auth.identity_source import KeycloakSource, identity_source_ref
|
||||
from litellm.llms.base_llm.auth.token_exchange import MAX_RESPONSE_BYTES
|
||||
|
||||
TOKEN_URL: Final = "https://keycloak.example/realms/litellm/protocol/openid-connect/token"
|
||||
CLIENT_ID: Final = "litellm"
|
||||
|
|
@ -331,3 +335,80 @@ class TestClientSecretNeverLeaks:
|
|||
)
|
||||
|
||||
assert CLIENT_SECRET not in caplog.text
|
||||
|
||||
|
||||
class StubHandler:
|
||||
"""Stands in for the HTTPHandler the default poster builds, so the poster's own contract
|
||||
(redirects off, error responses returned rather than raised, no-response guarded) is testable
|
||||
without a socket."""
|
||||
|
||||
def __init__(self, result: httpx.Response | Exception | None) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self._result = result
|
||||
|
||||
def post(self, url: str, *, content: bytes, headers: dict[str, str], timeout: float) -> httpx.Response | None:
|
||||
self.calls.append({"url": url, "content": content, "headers": headers, "timeout": timeout})
|
||||
if isinstance(self._result, Exception):
|
||||
raise self._result
|
||||
return self._result
|
||||
|
||||
|
||||
class TestDefaultKeycloakPoster:
|
||||
def test_builds_its_handler_once_with_redirects_disabled(self):
|
||||
built: list[StubHandler] = []
|
||||
|
||||
def factory() -> StubHandler:
|
||||
handler = StubHandler(httpx.Response(200, json={"access_token": "kc-token"}))
|
||||
built.append(handler)
|
||||
return handler
|
||||
|
||||
poster: Final = _HttpxSyncKeycloakPoster(handler_factory=factory) # pyright: ignore[reportArgumentType] # StubHandler stands in for the legacy-untyped HTTPHandler
|
||||
for _ in range(3):
|
||||
poster.post(TOKEN_URL, content=b"grant_type=client_credentials", headers={}, timeout=1.0)
|
||||
|
||||
assert len(built) == 1, "the handler is built once and reused"
|
||||
assert len(built[0].calls) == 3
|
||||
|
||||
def test_the_real_handler_refuses_to_follow_redirects(self):
|
||||
handler: Final = _new_keycloak_handler()
|
||||
assert handler.client.follow_redirects is False, (
|
||||
"a redirected token POST would replay the client secret to whatever host the redirect names"
|
||||
)
|
||||
|
||||
def test_an_http_status_error_becomes_its_response_rather_than_an_exception(self):
|
||||
response: Final = httpx.Response(
|
||||
401, json={"error": "invalid_client"}, request=httpx.Request("POST", TOKEN_URL)
|
||||
)
|
||||
poster: Final = _HttpxSyncKeycloakPoster(
|
||||
handler_factory=lambda: StubHandler(
|
||||
httpx.HTTPStatusError("boom", request=response.request, response=response)
|
||||
) # pyright: ignore[reportArgumentType] # StubHandler stands in for the legacy-untyped HTTPHandler
|
||||
)
|
||||
|
||||
assert poster.post(TOKEN_URL, content=b"", headers={}, timeout=1.0).status_code == 401
|
||||
|
||||
def test_a_missing_response_is_a_transport_error_not_a_none_deref(self):
|
||||
poster: Final = _HttpxSyncKeycloakPoster(handler_factory=lambda: StubHandler(None)) # pyright: ignore[reportArgumentType] # StubHandler stands in for the legacy-untyped HTTPHandler
|
||||
|
||||
with pytest.raises(httpx.TransportError):
|
||||
poster.post(TOKEN_URL, content=b"", headers={}, timeout=1.0)
|
||||
|
||||
|
||||
class TestDefaultSecretReader:
|
||||
def test_reads_through_litellm_secret_resolution(self, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("KEYCLOAK_CLIENT_SECRET_FOR_DEFAULT_READER", CLIENT_SECRET)
|
||||
|
||||
assert _default_secret_reader("os.environ/KEYCLOAK_CLIENT_SECRET_FOR_DEFAULT_READER") == CLIENT_SECRET
|
||||
|
||||
def test_an_unset_reference_reads_as_none_so_the_caller_raises(self):
|
||||
assert _default_secret_reader("os.environ/DEFINITELY_NOT_SET_KEYCLOAK_SECRET_REF") is None
|
||||
|
||||
|
||||
class TestOversizedSuccessBody:
|
||||
def test_a_success_body_over_the_cap_is_refused_before_it_is_parsed(self):
|
||||
oversized: Final = httpx.Response(200, content=b'{"access_token": "' + b"x" * MAX_RESPONSE_BYTES + b'"}')
|
||||
|
||||
with pytest.raises(ValueError, match="exceeded the size cap"):
|
||||
fetch_keycloak_assertion(
|
||||
make_config(), poster=ScriptedPoster([oversized]), secret_reader=DEFAULT_SECRET_READER
|
||||
)
|
||||
|
|
|
|||
|
|
@ -19,6 +19,10 @@ from litellm.llms.base_llm.auth.token_exchange import (
|
|||
MAX_ASSERTION_BYTES,
|
||||
MAX_RESPONSE_BYTES,
|
||||
JwtBearerTokenExchangeEngine,
|
||||
_default_assertion_reader,
|
||||
_error_summary,
|
||||
_HttpxSyncTokenPoster,
|
||||
_new_exchange_handler,
|
||||
redact_oauth_error_body,
|
||||
)
|
||||
from litellm.llms.base_llm.auth.types import (
|
||||
|
|
@ -35,6 +39,7 @@ from litellm.secret_managers.main import OidcPathNotAllowedError, _resolve_oidc_
|
|||
|
||||
DEFAULT_REF: Final = "oidc/env/TEST_ASSERTION"
|
||||
DEFAULT_ASSERTION: Final = "test-jwt-assertion"
|
||||
EXCHANGE_URL: Final = "https://token.example/v1/oauth/token"
|
||||
|
||||
|
||||
class FakeClock:
|
||||
|
|
@ -1118,3 +1123,100 @@ def test_invalidate_bypasses_lead_backoff():
|
|||
assert isinstance(second, MintedToken)
|
||||
assert second.access_token.get_secret_value() == "post-invalidate"
|
||||
assert len(poster.requests) == 2
|
||||
|
||||
|
||||
class StubExchangeHandler:
|
||||
"""Stands in for the HTTPHandler the default poster builds, so the poster's own contract is
|
||||
testable without a socket."""
|
||||
|
||||
def __init__(self, result: httpx.Response | Exception | None) -> None:
|
||||
self.calls = 0
|
||||
self._result = result
|
||||
|
||||
def post(self, url: str, *, content: bytes, headers: dict[str, str], timeout: float) -> httpx.Response | None:
|
||||
self.calls += 1
|
||||
if isinstance(self._result, Exception):
|
||||
raise self._result
|
||||
return self._result
|
||||
|
||||
|
||||
class TestDefaultTokenPoster:
|
||||
def test_builds_its_handler_once_and_reuses_it(self):
|
||||
built: list[StubExchangeHandler] = []
|
||||
|
||||
def factory() -> StubExchangeHandler:
|
||||
handler = StubExchangeHandler(httpx.Response(200, json={"access_token": "t"}))
|
||||
built.append(handler)
|
||||
return handler
|
||||
|
||||
poster: Final = _HttpxSyncTokenPoster(handler_factory=factory) # pyright: ignore[reportArgumentType] # StubExchangeHandler stands in for the legacy-untyped HTTPHandler
|
||||
for _ in range(3):
|
||||
poster.post(EXCHANGE_URL, content=b"", headers={}, timeout=1.0)
|
||||
|
||||
assert len(built) == 1
|
||||
assert built[0].calls == 3
|
||||
|
||||
def test_the_real_handler_refuses_to_follow_redirects(self):
|
||||
assert _new_exchange_handler().client.follow_redirects is False, (
|
||||
"a redirected exchange POST would replay the workload assertion to the redirect target"
|
||||
)
|
||||
|
||||
def test_an_http_status_error_becomes_its_response(self):
|
||||
response: Final = httpx.Response(
|
||||
401, json={"error": "invalid_grant"}, request=httpx.Request("POST", EXCHANGE_URL)
|
||||
)
|
||||
poster: Final = _HttpxSyncTokenPoster(
|
||||
handler_factory=lambda: StubExchangeHandler( # pyright: ignore[reportArgumentType] # StubExchangeHandler stands in for the legacy-untyped HTTPHandler
|
||||
httpx.HTTPStatusError("boom", request=response.request, response=response)
|
||||
)
|
||||
)
|
||||
|
||||
assert poster.post(EXCHANGE_URL, content=b"", headers={}, timeout=1.0).status_code == 401
|
||||
|
||||
def test_a_missing_response_is_a_transport_error(self):
|
||||
poster: Final = _HttpxSyncTokenPoster(handler_factory=lambda: StubExchangeHandler(None)) # pyright: ignore[reportArgumentType] # StubExchangeHandler stands in for the legacy-untyped HTTPHandler
|
||||
|
||||
with pytest.raises(httpx.TransportError):
|
||||
poster.post(EXCHANGE_URL, content=b"", headers={}, timeout=1.0)
|
||||
|
||||
|
||||
class TestDefaultAssertionReader:
|
||||
def test_reads_through_litellm_secret_resolution(self, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("WIF_ASSERTION_FOR_DEFAULT_READER", "header.payload.signature")
|
||||
|
||||
assert _default_assertion_reader("os.environ/WIF_ASSERTION_FOR_DEFAULT_READER") == "header.payload.signature"
|
||||
|
||||
def test_an_unset_reference_reads_as_none(self):
|
||||
assert _default_assertion_reader("os.environ/DEFINITELY_NOT_SET_WIF_ASSERTION_REF") is None
|
||||
|
||||
|
||||
class TestErrorSummary:
|
||||
def test_every_error_variant_summarises_without_carrying_a_secret(self):
|
||||
summaries: Final = {
|
||||
_error_summary(AssertionSourceError(kind="unreadable", source_ref="oidc/file/x")),
|
||||
_error_summary(InsecureTokenUrl(host="token.internal")),
|
||||
_error_summary(TokenEndpointError(status_code=401, redacted_body="invalid_grant")),
|
||||
_error_summary(TokenTransportError(detail="ConnectError: refused")),
|
||||
_error_summary(MalformedTokenResponse(detail="empty access_token")),
|
||||
}
|
||||
|
||||
assert {s.split(":")[0] for s in summaries} == {
|
||||
"AssertionSourceError",
|
||||
"InsecureTokenUrl",
|
||||
"TokenEndpointError",
|
||||
"TokenTransportError",
|
||||
"MalformedTokenResponse",
|
||||
}, "each variant names itself so a log line says which stage failed"
|
||||
|
||||
|
||||
class TestNonBearerTokenType:
|
||||
def test_a_non_bearer_token_type_is_refused(self):
|
||||
poster: Final = ScriptedPoster(
|
||||
[httpx.Response(200, json={"access_token": "tok", "token_type": "mac", "expires_in": 300})]
|
||||
)
|
||||
engine: Final = JwtBearerTokenExchangeEngine(poster=poster, assertion_reader=lambda _ref: DEFAULT_ASSERTION)
|
||||
|
||||
result: Final = engine.get_token(make_spec())
|
||||
|
||||
assert isinstance(result, MalformedTokenResponse)
|
||||
assert "non-bearer" in result.detail
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue