mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(spend-logs): read used_client_oauth_token from the bucket the route stamped
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
A guardrail on the unified path adds litellm_metadata to a chat request after the proxy stamped metadata, so both spend row writers read the new bucket and stored null. The success row now resolves the flag the same way the callback payload does, and the failure row picks the bucket from the request route.
This commit is contained in:
parent
c4f6d82a4f
commit
79bda61212
7 changed files with 101 additions and 17 deletions
|
|
@ -339,6 +339,13 @@ def get_or_create_metadata_bucket(
|
|||
return metadata_key, metadata_bucket
|
||||
|
||||
|
||||
def proxy_stamped_used_client_oauth_token(metadata: object, litellm_params: Mapping[str, object] | None) -> object:
|
||||
litellm_metadata: Final = litellm_params.get("litellm_metadata") if litellm_params is not None else None
|
||||
if isinstance(litellm_metadata, Mapping) and "used_client_oauth_token" in litellm_metadata:
|
||||
return litellm_metadata["used_client_oauth_token"]
|
||||
return metadata.get("used_client_oauth_token") if isinstance(metadata, Mapping) else None
|
||||
|
||||
|
||||
def get_litellm_metadata_from_kwargs(kwargs: dict):
|
||||
"""
|
||||
Helper to get litellm metadata from all litellm request kwargs
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ from litellm.litellm_core_utils.classifier_logging import (
|
|||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_provider_response_headers_from_hidden_params,
|
||||
is_expected_client_error,
|
||||
proxy_stamped_used_client_oauth_token,
|
||||
reconstruct_model_name,
|
||||
set_response_cost_in_hidden_params,
|
||||
)
|
||||
|
|
@ -288,13 +289,6 @@ _STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = (
|
|||
)
|
||||
|
||||
|
||||
def _proxy_stamped_used_client_oauth_token(metadata: object, litellm_params: Mapping[str, object] | None) -> object:
|
||||
litellm_metadata: Final = litellm_params.get("litellm_metadata") if litellm_params is not None else None
|
||||
if isinstance(litellm_metadata, Mapping) and "used_client_oauth_token" in litellm_metadata:
|
||||
return litellm_metadata["used_client_oauth_token"]
|
||||
return metadata.get("used_client_oauth_token") if isinstance(metadata, Mapping) else None
|
||||
|
||||
|
||||
def _get_provider_request_id(original_exception: Exception) -> str | None:
|
||||
try:
|
||||
error_response: Final = getattr(original_exception, "response", None)
|
||||
|
|
@ -5793,7 +5787,7 @@ class StandardLoggingPayloadSetup:
|
|||
team_alias=None,
|
||||
team_id=None,
|
||||
used_client_oauth_token=resolve_used_client_oauth_token(
|
||||
_proxy_stamped_used_client_oauth_token(metadata, litellm_params),
|
||||
proxy_stamped_used_client_oauth_token(metadata, litellm_params),
|
||||
custom_llm_provider,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ from litellm.proxy.db.db_spend_update_writer import (
|
|||
debitable_model_access_groups,
|
||||
get_llm_router,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, metadata_variable_name_for_route
|
||||
from litellm.proxy.spend_tracking.spend_event import (
|
||||
ObjectMapping,
|
||||
SpendEventBuildError,
|
||||
|
|
@ -86,8 +86,15 @@ _CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset(
|
|||
)
|
||||
|
||||
|
||||
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))
|
||||
def _proxy_stamped_used_client_oauth_token(
|
||||
request_data: Mapping[str, object], request_route: str | None
|
||||
) -> bool | None:
|
||||
proxy_bucket: Final = (
|
||||
get_metadata_variable_name_from_kwargs(request_data)
|
||||
if request_route is None
|
||||
else metadata_variable_name_for_route(request_route)
|
||||
)
|
||||
proxy_metadata: Final = request_data.get(proxy_bucket)
|
||||
stamped: Final = proxy_metadata.get("used_client_oauth_token") if isinstance(proxy_metadata, dict) else None
|
||||
return stamped if isinstance(stamped, bool) else None
|
||||
|
||||
|
|
@ -198,7 +205,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||
metadata=_metadata, original_exception=original_exception
|
||||
)
|
||||
|
||||
_metadata["used_client_oauth_token"] = _proxy_stamped_used_client_oauth_token(request_data)
|
||||
_metadata["used_client_oauth_token"] = _proxy_stamped_used_client_oauth_token(request_data, request_route)
|
||||
|
||||
existing_metadata: Final[dict] = request_data.get("metadata", None) or {}
|
||||
existing_metadata.update(_metadata)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from collections import OrderedDict
|
|||
from collections.abc import Mapping, MutableMapping, Sequence
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from pydantic import TypeAdapter
|
||||
|
|
@ -649,11 +649,14 @@ def _get_metadata_variable_name(request: Request) -> str:
|
|||
# Inline imports — auth_utils/route_checks participate in a proxy import cycle.
|
||||
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
|
||||
|
||||
path: Final = get_request_route(request)
|
||||
if "thread" in path or "assistant" in path:
|
||||
return metadata_variable_name_for_route(get_request_route(request))
|
||||
|
||||
|
||||
def metadata_variable_name_for_route(route: str) -> Literal["metadata", "litellm_metadata"]:
|
||||
if "thread" in route or "assistant" in route:
|
||||
return "litellm_metadata"
|
||||
|
||||
if any(route in path for route in LITELLM_METADATA_ROUTES):
|
||||
if any(metadata_route in route for metadata_route in LITELLM_METADATA_ROUTES):
|
||||
return "litellm_metadata"
|
||||
|
||||
return "metadata"
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ from litellm.constants import (
|
|||
from litellm.litellm_core_utils.classifier_logging import classifier_audit_fields, without_classifier_audit
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_litellm_metadata_from_kwargs,
|
||||
proxy_stamped_used_client_oauth_token,
|
||||
reconstruct_model_name,
|
||||
)
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
|
||||
|
|
@ -720,7 +721,7 @@ def get_logging_payload(
|
|||
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
|
||||
proxy_stamped_used_client_oauth_token(litellm_params.get("metadata"), litellm_params), custom_llm_provider
|
||||
),
|
||||
azure_spillover=azure_spillover(
|
||||
response_headers=kwargs.get("response_headers")
|
||||
|
|
|
|||
|
|
@ -245,6 +245,53 @@ async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_
|
|||
assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"request_route, metadata_buckets, expected",
|
||||
[
|
||||
(
|
||||
"/v1/chat/completions",
|
||||
{"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_api_key_hash": "guardrail"}},
|
||||
True,
|
||||
),
|
||||
(
|
||||
"/v1/messages",
|
||||
{"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_api_key_hash": "proxy"}},
|
||||
None,
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_async_post_call_failure_hook_reads_used_client_oauth_token_from_the_routes_stamped_bucket(
|
||||
request_route: str, metadata_buckets: dict, expected: bool | None
|
||||
):
|
||||
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", request_route=request_route),
|
||||
)
|
||||
|
||||
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
|
||||
|
|
|
|||
|
|
@ -3305,6 +3305,31 @@ def test_get_logging_payload_records_used_client_oauth_token_for_the_selected_pr
|
|||
assert _get_spend_logs_metadata(None)["used_client_oauth_token"] is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"litellm_params, expected",
|
||||
[
|
||||
(
|
||||
{"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_api_key_hash": "guardrail"}},
|
||||
True,
|
||||
),
|
||||
(
|
||||
{"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"used_client_oauth_token": False}},
|
||||
False,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_get_logging_payload_reads_used_client_oauth_token_from_the_bucket_the_proxy_stamped(
|
||||
litellm_params: dict, expected: bool
|
||||
):
|
||||
payload = get_logging_payload(
|
||||
kwargs={"model": "claude-sonnet-5", "custom_llm_provider": "anthropic", "litellm_params": litellm_params},
|
||||
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
|
||||
|
||||
|
||||
def test_redact_logged_api_key_bearer_only_returns_none():
|
||||
# "bearer " with nothing after stripping is equivalent to no key
|
||||
assert _redact_logged_api_key("bearer ") is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue