fix(proxy): keep the proxy's used_client_oauth_token stamp on failure rows and move the resolver under llms/anthropic

This commit is contained in:
mateo-berri 2026-09-24 17:41:48 -07:00
parent 05b35191d3
commit 2d784337b3
8 changed files with 84 additions and 29 deletions

View file

@ -1,15 +1,7 @@
from collections.abc import Sequence
from typing import Final
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
from litellm.types.utils import ProviderSpecificHeader
class ProviderSpecificHeaderUtils:

View file

@ -79,7 +79,6 @@ 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,
@ -5745,6 +5744,9 @@ class StandardLoggingPayloadSetup:
- If the input metadata is None or not a dictionary, an empty StandardLoggingMetadata object is returned.
- If 'user_api_key' is present in metadata and is a valid SHA256 hash, it's stored as 'user_api_key_hash'.
"""
from litellm.llms.anthropic.common_utils import ( # noqa: PLC0415 # that module imports this one transitively
resolve_used_client_oauth_token,
)
prompt_management_metadata: StandardLoggingPromptManagementMetadata | None = None
if litellm_params is not None:

View file

@ -39,6 +39,7 @@ from litellm.types.llms.anthropic import (
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.model_listing import ModelInfoResponse
from litellm.types.utils import LlmProviders
_MessageT = TypeVar("_MessageT")
@ -225,6 +226,15 @@ def is_anthropic_oauth_key(value: str | None) -> bool:
return value.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX)
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
def _merge_beta_headers(existing: str | None, new_beta: str) -> str:
"""Merge a new beta value into an existing comma-separated anthropic-beta header."""
if not existing:

View file

@ -13,6 +13,7 @@ from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
budget_reservation_from_metadata,
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
)
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost
@ -83,10 +84,12 @@ _CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset(
str(CallTypes.aretrieve_batch),
)
)
_FAILURE_ROW_KEYS_LIFTED_FROM_LITELLM_METADATA: Final[tuple[str, ...]] = (
"standard_logging_guardrail_information",
"used_client_oauth_token",
)
def _proxy_stamped_used_client_oauth_token(request_data: Mapping[str, object]) -> bool | None:
proxy_metadata: Final = request_data.get(get_metadata_variable_name_from_kwargs(request_data))
stamped: Final = proxy_metadata.get("used_client_oauth_token") if isinstance(proxy_metadata, dict) else None
return stamped if isinstance(stamped, bool) else None
def _proxy_spend_writer() -> DBSpendUpdateWriter:
@ -195,17 +198,19 @@ class _ProxyDBLogger(CustomLogger):
metadata=_metadata, original_exception=original_exception
)
_metadata["used_client_oauth_token"] = _proxy_stamped_used_client_oauth_token(request_data)
existing_metadata: Final[dict] = request_data.get("metadata", None) or {}
existing_metadata.update(_metadata)
litellm_metadata_bucket: Final = request_data.get("litellm_metadata")
existing_metadata.update(
(key, litellm_metadata_bucket[key])
for key in _FAILURE_ROW_KEYS_LIFTED_FROM_LITELLM_METADATA
if isinstance(litellm_metadata_bucket, dict)
and key not in existing_metadata
and litellm_metadata_bucket.get(key) is not None
)
if (
isinstance(litellm_metadata_bucket, dict)
and "standard_logging_guardrail_information" not in existing_metadata
):
guardrail_info: Final = litellm_metadata_bucket.get("standard_logging_guardrail_information")
if guardrail_info is not None:
existing_metadata["standard_logging_guardrail_information"] = guardrail_info
if "litellm_params" not in request_data:
request_data["litellm_params"] = {}

View file

@ -34,7 +34,6 @@ 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,
@ -46,6 +45,7 @@ from litellm.litellm_core_utils.url_utils import (
is_url_destination_allowed_by_host,
provider_url_destination_candidates,
)
from litellm.llms.anthropic.common_utils import ANTHROPIC_OAUTH_FORWARD_PROVIDERS
from litellm.proxy._types import (
AddTeamCallback,
CommonProxyErrors,

View file

@ -2429,13 +2429,15 @@ async def ui_view_spend_logs(
default=None,
description="Filter logs by cache state: 'hit' or 'miss'. Miss includes legacy rows with a null/unknown cache state",
),
used_client_oauth_token: bool | None = fastapi.Query(
default=None,
description=(
"Filter logs by the credential the upstream call used: true for a client-forwarded Anthropic OAuth token, "
"false for the deployment's configured key. Rows written before this flag existed match neither"
used_client_oauth_token: Annotated[
bool | None,
fastapi.Query(
description=(
"Filter logs by the credential the upstream call used: true for a client-forwarded Anthropic OAuth "
"token, false for the deployment's configured key. Rows written before this flag existed match neither"
),
),
),
] = None,
span_type: str | None = fastapi.Query(
default=None,
description="Filter logs by span type: llm, agent, mcp, or batch",

View file

@ -35,7 +35,6 @@ 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,
@ -44,6 +43,7 @@ from litellm.litellm_core_utils.litellm_logging import (
)
from litellm.litellm_core_utils.ptu_pricing import azure_spillover
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
from litellm.llms.anthropic.common_utils import resolve_used_client_oauth_token
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error

View file

@ -201,6 +201,50 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from
assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected
@pytest.mark.asyncio
@pytest.mark.parametrize(
"metadata_buckets, expected",
[
({"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"used_client_oauth_token": False}}, False),
({"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_id": "caller"}}, None),
({"metadata": {"used_client_oauth_token": "yes"}}, None),
],
)
async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_client_oauth_token(
metadata_buckets: dict, expected: bool | None
):
"""
On /v1/messages and /v1/responses the request's own metadata field belongs to the caller, so a
used_client_oauth_token they put there must never outrank the proxy's stamp or stand in for a missing one
"""
logger = _ProxyDBLogger()
request_data = {
"model": "claude-sonnet-5",
"custom_llm_provider": "anthropic",
"messages": [{"role": "user", "content": "Hello"}],
"proxy_server_request": {"request_id": "test_request_id"},
**metadata_buckets,
}
with patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database:
await logger.async_post_call_failure_hook(
request_data=request_data,
original_exception=Exception("rate limited"),
user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"),
)
payload = get_logging_payload(
kwargs=mock_update_database.call_args[1]["kwargs"],
response_obj={},
start_time=datetime.now(),
end_time=datetime.now(),
)
assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_request():
"""LIT-5651: a request blocked by a guardrail never reaches the LLM, but the