fix(proxy): resolve used_client_oauth_token against the provider the call was sent to

This commit is contained in:
mateo-berri 2026-09-24 16:55:11 -07:00
parent d688293c81
commit 05b35191d3
8 changed files with 114 additions and 16 deletions

View file

@ -1,7 +1,15 @@
from collections.abc import Sequence
from typing import Final
from litellm.types.utils import ProviderSpecificHeader
from litellm.types.utils import LlmProviders, ProviderSpecificHeader
ANTHROPIC_OAUTH_FORWARD_PROVIDERS: Final[frozenset[str]] = frozenset((LlmProviders.ANTHROPIC.value,))
def resolve_used_client_oauth_token(client_sent_oauth_token: object, custom_llm_provider: str | None) -> bool | None:
if not isinstance(client_sent_oauth_token, bool):
return None
return client_sent_oauth_token and custom_llm_provider in ANTHROPIC_OAUTH_FORWARD_PROVIDERS
class ProviderSpecificHeaderUtils:

View file

@ -79,6 +79,7 @@ from litellm.litellm_core_utils.core_helpers import (
)
from litellm.litellm_core_utils.error_normalization import normalize_error
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.get_provider_specific_headers import resolve_used_client_oauth_token
from litellm.litellm_core_utils.internal_call_metadata import (
MODEL_ACCESS_GROUP_METADATA_KEY,
is_unbilled_non_inference_call,
@ -283,7 +284,10 @@ else:
_PAGERDUTY_ALERTING_FACTORY: Final = PagerDutyAlerting
_in_memory_loggers: Final[list[CustomLogger]] = []
_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = frozenset(StandardLoggingMetadata.__annotations__.keys())
_STANDARD_LOGGING_METADATA_RESOLVED_KEYS: Final[frozenset[str]] = frozenset(("used_client_oauth_token",))
_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = (
frozenset(StandardLoggingMetadata.__annotations__.keys()) - _STANDARD_LOGGING_METADATA_RESOLVED_KEYS
)
def _get_provider_request_id(original_exception: Exception) -> str | None:
@ -5726,6 +5730,7 @@ class StandardLoggingPayloadSetup:
proxy_server_request: dict | None = None,
start_time: dt_object | None = None,
response_id: str | None = None,
custom_llm_provider: str | None = None,
) -> StandardLoggingMetadata:
"""
Clean and filter the metadata dictionary to include only the specified keys in StandardLoggingMetadata.
@ -5789,7 +5794,10 @@ class StandardLoggingPayloadSetup:
user_api_key_auth_metadata=None,
team_alias=None,
team_id=None,
used_client_oauth_token=None,
used_client_oauth_token=resolve_used_client_oauth_token(
metadata.get("used_client_oauth_token") if isinstance(metadata, dict) else None,
custom_llm_provider,
),
)
if isinstance(metadata, dict):
for key in metadata.keys() & _STANDARD_LOGGING_METADATA_KEYS:
@ -6513,6 +6521,7 @@ def get_standard_logging_object_payload(
stream=kwargs.get("stream", False),
)
# clean up litellm metadata
selected_provider: Final = kwargs.get("custom_llm_provider")
clean_metadata: Final = StandardLoggingPayloadSetup.get_standard_logging_metadata(
metadata=metadata,
litellm_params=litellm_params,
@ -6524,6 +6533,7 @@ def get_standard_logging_object_payload(
proxy_server_request=proxy_server_request,
start_time=start_time,
response_id=id,
custom_llm_provider=selected_provider if isinstance(selected_provider, str) else None,
)
_request_body: Final = proxy_server_request.get("body", {})
end_user_id: Final = clean_metadata["user_api_key_end_user_id"] or _request_body.get(

View file

@ -34,6 +34,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.core_helpers import is_codex_user_agent
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.get_provider_specific_headers import ANTHROPIC_OAUTH_FORWARD_PROVIDERS
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
TRUSTED_CALLBACK_VARS_FIELD,
_request_blocked_callback_params,
@ -3456,7 +3457,7 @@ _ANTHROPIC_API_HEADER_PROVIDERS: Final = ",".join(
LlmProviders.VERTEX_AI.value,
)
)
_ANTHROPIC_OAUTH_CREDENTIAL_PROVIDERS: Final = LlmProviders.ANTHROPIC.value
_ANTHROPIC_OAUTH_CREDENTIAL_PROVIDERS: Final = ",".join(sorted(ANTHROPIC_OAUTH_FORWARD_PROVIDERS))
def add_provider_specific_headers_to_request(

View file

@ -35,6 +35,7 @@ from litellm.litellm_core_utils.core_helpers import (
reconstruct_model_name,
)
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
from litellm.litellm_core_utils.get_provider_specific_headers import resolve_used_client_oauth_token
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call
from litellm.litellm_core_utils.litellm_logging import (
coerce_model_access_groups,
@ -144,6 +145,7 @@ _STAMPED_METADATA_KEYS: Final = frozenset(
"autorouter_savings",
"autorouter_savings_estimate",
"autorouter_baseline_observation",
"used_client_oauth_token",
)
)
@ -168,6 +170,7 @@ def _get_spend_logs_metadata(
autorouter_baseline_observation: str | None = None,
router_metadata: SpendLogsRouterMetadata | None = None,
azure_spillover: AzureSpillover | None = None,
used_client_oauth_token: bool | None = None,
) -> SpendLogsMetadata:
if metadata is None:
return SpendLogsMetadata(
@ -212,7 +215,7 @@ def _get_spend_logs_metadata(
litellm_call_id=litellm_call_id,
router_metadata=router_metadata,
azure_spillover=azure_spillover,
used_client_oauth_token=None,
used_client_oauth_token=used_client_oauth_token,
)
verbose_proxy_logger.debug(
"getting payload for SpendLogs, available keys in metadata: " + str(list(metadata.keys()))
@ -228,6 +231,7 @@ def _get_spend_logs_metadata(
autorouter_baseline_observation=autorouter_baseline_observation,
router_metadata=router_metadata,
azure_spillover=azure_spillover,
used_client_oauth_token=used_client_oauth_token,
)
_raw_key: Final = clean_metadata.get("user_api_key")
_trusted_hash: Final = metadata.get("user_api_key_hash")
@ -695,6 +699,9 @@ def get_logging_payload(
selected_provider=custom_llm_provider,
router_correlation_id=litellm_call_id,
),
used_client_oauth_token=resolve_used_client_oauth_token(
metadata.get("used_client_oauth_token") if metadata is not None else None, custom_llm_provider
),
azure_spillover=azure_spillover(
response_headers=kwargs.get("response_headers")
if isinstance(kwargs.get("response_headers"), Mapping)

View file

@ -4846,6 +4846,39 @@ def test_get_standard_logging_object_payload_includes_litellm_call_id(logging_ob
assert payload["litellm_call_id"] == call_id
@pytest.mark.parametrize(
"client_sent_oauth_token, custom_llm_provider, expected",
[(True, "anthropic", True), (True, "bedrock", False), (False, "anthropic", False), (None, "anthropic", None)],
)
def test_get_standard_logging_object_payload_resolves_used_client_oauth_token_against_the_selected_provider(
logging_obj, client_sent_oauth_token: bool | None, custom_llm_provider: str, expected: bool | None
):
"""The proxy stamps whether the client presented an Anthropic OAuth bearer before routing, but the
bearer only reaches an Anthropic deployment, so the logged flag must follow the provider that was called."""
from datetime import datetime
from litellm.litellm_core_utils.litellm_logging import get_standard_logging_object_payload
request_metadata = {} if client_sent_oauth_token is None else {"used_client_oauth_token": client_sent_oauth_token}
now = datetime.now()
payload = get_standard_logging_object_payload(
kwargs={
"model": "claude-sonnet-5",
"messages": [],
"custom_llm_provider": custom_llm_provider,
"litellm_params": {"metadata": request_metadata},
},
init_response_obj={},
start_time=now,
end_time=now,
logging_obj=logging_obj,
status="success",
)
assert payload is not None
assert payload["metadata"]["used_client_oauth_token"] is expected
def test_get_standard_logging_object_payload_carries_matched_access_groups(logging_obj):
"""Access groups stamped at auth time reach the logging payload, so integrations see what a request billed."""
from datetime import datetime

View file

@ -161,9 +161,12 @@ async def test_async_post_call_failure_hook_does_not_clobber_guardrail_info_in_m
@pytest.mark.asyncio
@pytest.mark.parametrize("used_client_oauth_token", [True, False])
@pytest.mark.parametrize(
"used_client_oauth_token, custom_llm_provider, expected",
[(True, "anthropic", True), (True, "bedrock", False), (False, "anthropic", False)],
)
async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from_litellm_metadata(
used_client_oauth_token: bool,
used_client_oauth_token: bool, custom_llm_provider: str, expected: bool
):
"""
/v1/messages and /v1/responses stamp the proxy's own fields into request_data["litellm_metadata"]
@ -173,6 +176,7 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from
logger = _ProxyDBLogger()
request_data = {
"model": "claude-sonnet-5",
"custom_llm_provider": custom_llm_provider,
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {"user_id": "anthropic-native-metadata"},
"litellm_metadata": {"used_client_oauth_token": used_client_oauth_token},
@ -194,7 +198,7 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from
payload = get_logging_payload(
kwargs=call_kwargs, response_obj={}, start_time=datetime.now(), end_time=datetime.now()
)
assert json.loads(payload["metadata"])["used_client_oauth_token"] is used_client_oauth_token
assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected
@pytest.mark.asyncio

View file

@ -3271,12 +3271,38 @@ def test_get_spend_logs_metadata_keeps_user_agent():
assert _get_spend_logs_metadata(None)["user_agent"] is None
@pytest.mark.parametrize("used_client_oauth_token", [True, False])
def test_get_spend_logs_metadata_keeps_used_client_oauth_token(used_client_oauth_token: bool):
meta = _get_spend_logs_metadata({"used_client_oauth_token": used_client_oauth_token})
assert meta["used_client_oauth_token"] is used_client_oauth_token
@pytest.mark.parametrize(
"client_sent_oauth_token, custom_llm_provider, expected",
[
(True, "anthropic", True),
(True, "bedrock", False),
(True, "vertex_ai", False),
(False, "anthropic", False),
(None, "anthropic", None),
],
)
def test_get_logging_payload_records_used_client_oauth_token_for_the_selected_provider(
client_sent_oauth_token: bool | None, custom_llm_provider: str, expected: bool | None
):
"""The client's OAuth bearer is only forwarded to an Anthropic deployment, so a request that
the router sent to Bedrock or Vertex paid with the configured key and must not read true."""
request_metadata = (
{"user_agent": "claude-cli/2.1.0"}
if client_sent_oauth_token is None
else {"user_agent": "claude-cli/2.1.0", "used_client_oauth_token": client_sent_oauth_token}
)
payload = get_logging_payload(
kwargs={
"model": "claude-sonnet-5",
"custom_llm_provider": custom_llm_provider,
"litellm_params": {"metadata": request_metadata},
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected
assert _get_spend_logs_metadata(None)["used_client_oauth_token"] is None
assert _get_spend_logs_metadata({"user_agent": "curl/8.7.1"})["used_client_oauth_token"] is None
def test_redact_logged_api_key_bearer_only_returns_none():

View file

@ -38,8 +38,7 @@ from litellm.proxy.litellm_pre_call_utils import (
move_guardrails_to_metadata,
)
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
from litellm.litellm_core_utils.litellm_logging import get_standard_logging_metadata
from litellm.proxy.spend_tracking.spend_tracking_utils import _get_spend_logs_metadata
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY
from litellm.litellm_core_utils.redact_messages import _get_turn_off_message_logging_from_dynamic_params
from litellm.litellm_core_utils.get_provider_specific_headers import (
@ -6794,7 +6793,17 @@ async def test_add_litellm_data_to_request_stamps_used_client_oauth_token(path,
return updated[metadata_variable_name]
def spend_log_row_metadata(request_metadata: dict) -> dict:
return dict(_get_spend_logs_metadata(dict(get_standard_logging_metadata(metadata=request_metadata))))
row = get_logging_payload(
kwargs={
"model": "claude-sonnet-5",
"custom_llm_provider": "anthropic",
"litellm_params": {"metadata": request_metadata},
},
response_obj={},
start_time=datetime.now(timezone.utc),
end_time=datetime.now(timezone.utc),
)
return json.loads(row["metadata"])
seat_row = spend_log_row_metadata(
await metadata_for({"Authorization": _OAUTH_TOKEN, "x-litellm-api-key": "Bearer sk-virtual-key"})