mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
05b35191d3
commit
2d784337b3
8 changed files with 84 additions and 29 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"] = {}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue