fix(auth): compare both sides of a reflected credential, and cover the encoded forms

Reducing only the response to credential characters, while comparing it against an
untouched secret, stopped matching any secret carrying spaces or punctuation of its
own. A hand-set passphrase echoed back whole therefore reached the caller, which the
earlier contiguous match had caught. Both sides are reduced now, and a verbatim check
runs first so the result does not depend on what the secret is made of.

Percent-escaping is reversible and applies to any field, so the comparison also runs
over a decoded copy rather than asking each caller to enumerate that shape. What a
caller still declares is an encoding the redactor cannot reverse: client_secret_basic
sends base64 of id:secret, which decodes straight back to the secret.

The docstring no longer claims more than this does. Reducing to a shared alphabet
stops an accidental or naive echo; an endpoint that deliberately re-encodes or
interleaves the credential defeats any substring match, and this was never the control
keeping the credential from an endpoint that already holds it.
This commit is contained in:
derhornspieler 2026-08-25 07:55:22 -04:00
parent e14811bdbf
commit 7076219284
4 changed files with 84 additions and 25 deletions

View file

@ -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,
)

View file

@ -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:

View file

@ -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:

View file

@ -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."""