diff --git a/litellm/llms/anthropic/count_tokens/token_counter.py b/litellm/llms/anthropic/count_tokens/token_counter.py index a9a4bf0358e..2942033f277 100644 --- a/litellm/llms/anthropic/count_tokens/token_counter.py +++ b/litellm/llms/anthropic/count_tokens/token_counter.py @@ -61,9 +61,8 @@ class AnthropicTokenCounter(BaseTokenCounter): # None and the caller silently falls back to the local tokenizer, so a workload identity # deployment would never reach Anthropic's authoritative count. The minted token is an # sk-ant-oat, which get_required_headers already sends as a Bearer rather than x-api-key. - api_key: Final = static_key or await aget_anthropic_wif_token( - litellm_params, litellm_params.get("api_base"), model_to_use - ) + api_base: Final = litellm_params.get("api_base") + api_key: Final = static_key or await aget_anthropic_wif_token(litellm_params, api_base, model_to_use) if not api_key: verbose_logger.warning("No Anthropic credential found for token counting") @@ -74,6 +73,8 @@ class AnthropicTokenCounter(BaseTokenCounter): model=model_to_use, messages=messages, api_key=api_key, + # The token is minted for this base, so the count has to be asked of the same host. + api_base=api_base, tools=tools, system=system, ) diff --git a/litellm/llms/base_llm/auth/client_credentials.py b/litellm/llms/base_llm/auth/client_credentials.py index ec312bc63bd..f96feb0a1d4 100644 --- a/litellm/llms/base_llm/auth/client_credentials.py +++ b/litellm/llms/base_llm/auth/client_credentials.py @@ -148,14 +148,21 @@ def _resolve_client_secret(config: KeycloakSource, secret_reader: SecretReader) def _wire_forms_of_secret(config: KeycloakSource, client_secret: str) -> tuple[SecretStr, ...]: """Every shape the secret leaves this process in, so an echo of any of them is caught. - client_secret_basic sends base64 of ``id:secret``, which decodes straight back to the secret, - so an endpoint echoing that blob hands over reversible material that a raw comparison misses. + Neither grant sends the secret verbatim. client_secret_basic base64s ``id:secret``, which + decodes straight back to it, and client_secret_post percent-escapes it. An endpoint echoing + either shape hands over reversible material a raw comparison would miss. """ raw: Final = SecretStr(client_secret) - if config.auth_method != "client_secret_basic": - return (raw,) - encoded_pair: Final = f"{_form_encode(config.client_id)}:{_form_encode(client_secret)}" - return (raw, SecretStr(base64.b64encode(encoded_pair.encode()).decode("ascii"))) + match config.auth_method: + case "client_secret_basic": + 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=""))) + case _: + assert_never(config.auth_method) def _endpoint_error_message(config: KeycloakSource, response: httpx.Response, client_secret: str) -> str: diff --git a/litellm/llms/base_llm/auth/token_exchange.py b/litellm/llms/base_llm/auth/token_exchange.py index e0cab3f4ad2..6243a97a51a 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 urlencode, urlsplit, urlunsplit +from urllib.parse import unquote, urlencode, urlsplit, urlunsplit import httpx from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError @@ -141,28 +141,43 @@ def redact_oauth_error_body( def _drop_reflected_assertion(rendered: str, assertion: SecretStr | None) -> str: - """A token endpoint that echoes the submitted assertion back would otherwise put it in the - operator log and in the error handed to the caller. + """Catches an endpoint that echoes the submitted credential back, verbatim or in fragments, + however it split or percent-encoded it. - Matching a contiguous slice is not enough. An endpoint that returns a short piece of the - assertion, or returns it broken up by delimiters, shares no long run with it, so the credential - material would travel on unredacted. The rendered text is therefore stripped down to the - characters a credential is made of and scanned against the assertion, which catches a fragment - wherever it starts and however it was split. Over-redacting an error that merely happens to - contain such a run is the right way to be wrong here. + Both sides are reduced to the characters a credential is made of before comparison. Stripping + only the rendered side would stop matching a secret that carries spaces or punctuation of its + own, which is exactly the hand-set passphrase most at risk of being echoed. + + This stops an accidental or naive echo. It cannot stop an endpoint that deliberately re-encodes + or interleaves the credential, and it is not what keeps the credential from the endpoint, which + already holds it. What it protects is blast radius: keeping the value out of the caller's error + and out of third-party log sinks. """ if assertion is None: return rendered secret: Final = assertion.get_secret_value() if not secret: return rendered - if len(secret) <= _REFLECTION_MIN_RUN: - return _REFLECTED_VALUE_MESSAGE if secret in rendered else rendered - compacted: Final = _CREDENTIAL_CHARS.sub("", rendered) - runs: Final = ( - compacted[start : start + _REFLECTION_MIN_RUN] for start in range(len(compacted) - _REFLECTION_MIN_RUN + 1) - ) - return _REFLECTED_VALUE_MESSAGE if any(run in secret for run in runs) else rendered + if secret in rendered: + return _REFLECTED_VALUE_MESSAGE + compacted_secret: Final = _CREDENTIAL_CHARS.sub("", secret) + if not compacted_secret: + return rendered + return _REFLECTED_VALUE_MESSAGE if _shares_a_credential_run(rendered, compacted_secret) else rendered + + +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.""" + for candidate in (rendered, unquote(rendered)): + compacted: Final = _CREDENTIAL_CHARS.sub("", candidate) + if any( + compacted[start : start + _REFLECTION_MIN_RUN] in compacted_secret + for start in range(len(compacted) - _REFLECTION_MIN_RUN + 1) + ): + return True + return False def _redact_body_text(body_text: str) -> str: 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 b0e0929c5cc..4ff8858bea5 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 @@ -1,6 +1,8 @@ import asyncio import concurrent.futures +import base64 import json +from urllib.parse import quote import logging import threading import time @@ -603,6 +605,40 @@ class TestRedactionAndCaps: assert "PAYLOADMIDDLE" not in result.redacted_body assert tail[:40] not in result.redacted_body + def test_a_secret_carrying_spaces_is_dropped_when_echoed_whole(self): + """Regression on the redactor itself: comparing a compacted response against an + uncompacted secret stopped matching hand-set passphrases, which are exactly the secrets + most likely to be echoed and the ones an earlier contiguous match had caught.""" + assertion = SecretStr("correct horse battery staple, 42!") + echoed = assertion.get_secret_value() + body = {"error": "invalid_client", "error_description": f"secret {echoed} rejected"} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert echoed not in result.redacted_body + + def test_a_percent_encoded_secret_is_dropped(self): + """A form-encoded grant puts the secret on the wire percent-escaped, so an echo of that + shape has to be recognised without every caller enumerating it.""" + assertion = SecretStr("sUp3r+S3cret/Value=123") + echoed = quote(assertion.get_secret_value(), safe="") + body = {"error": "invalid_client", "error_description": f"rejected {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.""" + raw = SecretStr("sUp3rS3cretValue123") + blob = SecretStr(base64.b64encode(b"litellm:sUp3rS3cretValue123").decode()) + body = {"error": "invalid_client", "error_description": f"bad {blob.get_secret_value()}"} + + result = redact_oauth_error_body(400, json.dumps(body), (raw, blob)) + + assert blob.get_secret_value() not in result.redacted_body + def test_a_fragment_shorter_than_a_long_run_is_dropped(self): """A slice too short to share a long contiguous run with the assertion is still assertion material, and repeated errors would hand it over piece by piece."""