diff --git a/litellm/llms/base_llm/auth/client_credentials.py b/litellm/llms/base_llm/auth/client_credentials.py index f96feb0a1d4..4195ec3ea2a 100644 --- a/litellm/llms/base_llm/auth/client_credentials.py +++ b/litellm/llms/base_llm/auth/client_credentials.py @@ -15,7 +15,7 @@ import threading from collections.abc import Callable, Mapping from types import MappingProxyType from typing import TYPE_CHECKING, Final, TypeAlias -from urllib.parse import quote, urlencode +from urllib.parse import quote, quote_plus, urlencode import httpx from pydantic import BaseModel, SecretStr, ValidationError @@ -158,9 +158,10 @@ def _wire_forms_of_secret(config: KeycloakSource, client_secret: str) -> tuple[S encoded_pair: Final = f"{_form_encode(config.client_id)}:{_form_encode(client_secret)}" return (raw, SecretStr(base64.b64encode(encoded_pair.encode()).decode("ascii"))) case "client_secret_post": - # urlencode percent-escapes reserved characters, so a secret containing any of them - # leaves in a shape the raw comparison would not recognise coming back. - return (raw, SecretStr(quote(client_secret, safe=""))) + # urlencode escapes reserved characters and writes a space as "+", so a secret + # containing either leaves in a shape the raw comparison would not recognise coming + # back. quote_plus is what urlencode itself applies. + return (raw, SecretStr(quote_plus(client_secret))) case _: assert_never(config.auth_method) diff --git a/litellm/llms/base_llm/auth/token_exchange.py b/litellm/llms/base_llm/auth/token_exchange.py index a6567455490..9da2a942b9c 100644 --- a/litellm/llms/base_llm/auth/token_exchange.py +++ b/litellm/llms/base_llm/auth/token_exchange.py @@ -19,7 +19,7 @@ from dataclasses import dataclass from math import inf from types import MappingProxyType from typing import TYPE_CHECKING, Final, Protocol, TypeAlias -from urllib.parse import unquote, urlencode, urlsplit, urlunsplit +from urllib.parse import unquote, unquote_plus, urlencode, urlsplit, urlunsplit import httpx from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError @@ -170,8 +170,11 @@ def _shares_a_credential_run(rendered: str, compacted_secret: str) -> bool: """``unquote`` covers a credential sent form-encoded, without every caller enumerating that shape for itself: percent-escaping is reversible and applies to any field, query string included.""" + # unquote covers %XX; unquote_plus additionally covers the "+" a form-encoded body uses for a + # space. Both are kept rather than only the wider one, because "+" is a base64 character and + # decoding it away would lose a run that the undecoded candidate still matches on. compacted_candidates: Final = tuple( - _CREDENTIAL_CHARS.sub("", candidate) for candidate in (rendered, unquote(rendered)) + _CREDENTIAL_CHARS.sub("", candidate) for candidate in (rendered, unquote(rendered), unquote_plus(rendered)) ) return any( compacted[start : start + _REFLECTION_MIN_RUN] in compacted_secret diff --git a/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py b/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py index 4ff8858bea5..5e0be11e0d6 100644 --- a/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py +++ b/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py @@ -2,7 +2,7 @@ import asyncio import concurrent.futures import base64 import json -from urllib.parse import quote +from urllib.parse import quote, urlencode import logging import threading import time @@ -628,6 +628,18 @@ class TestRedactionAndCaps: assert echoed not in result.redacted_body + def test_a_space_encoded_as_plus_is_dropped(self): + """A form-encoded body writes a space as "+", not %20, so percent-decoding alone does not + recover the secret and a passphrase echoed in its wire shape would travel on.""" + assertion = SecretStr("correct horse battery staple") + echoed = urlencode({"client_secret": assertion.get_secret_value()}).split("=", 1)[1] + body = {"error": "invalid_client", "error_description": f"rejected {echoed}"} + + assert "+" in echoed + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert echoed not in result.redacted_body + def test_several_wire_forms_are_all_compared(self): """The caller declares each shape it sent, since an encoding the redactor cannot reverse (base64 of id:secret) is only knowable there."""