Skip to content
Merged
7 changes: 7 additions & 0 deletions litellm/litellm_core_utils/core_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
17 changes: 16 additions & 1 deletion litellm/litellm_core_utils/litellm_logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,7 @@
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,
)
Expand Down Expand Up @@ -284,7 +285,10 @@
_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:
Expand Down Expand Up @@ -5730,6 +5734,7 @@ def get_standard_logging_metadata(
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.
Expand All @@ -5744,6 +5749,9 @@ def get_standard_logging_metadata(
- 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:
Expand Down Expand Up @@ -5793,6 +5801,10 @@ def get_standard_logging_metadata(
user_api_key_auth_metadata=None,
team_alias=None,
team_id=None,
used_client_oauth_token=resolve_used_client_oauth_token(
proxy_stamped_used_client_oauth_token(metadata, litellm_params),
custom_llm_provider,
),
Comment thread
cursor[bot] marked this conversation as resolved.
)
if isinstance(metadata, dict):
for key in metadata.keys() & _STANDARD_LOGGING_METADATA_KEYS:
Expand Down Expand Up @@ -6516,6 +6528,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,
Expand All @@ -6527,6 +6540,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(
Expand Down Expand Up @@ -6801,6 +6815,7 @@ def get_standard_logging_metadata(
user_api_key_auth_metadata=None,
team_alias=None,
team_id=None,
used_client_oauth_token=None,
)
if isinstance(metadata, dict):
# Update the clean_metadata with values from input metadata that match StandardLoggingMetadata fields
Expand Down
10 changes: 10 additions & 0 deletions litellm/llms/anthropic/common_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.model_listing import ModelInfoResponse
from litellm.types.utils import LlmProviders

_MessageT = TypeVar("_MessageT")

Expand Down Expand Up @@ -226,6 +227,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:
Expand Down
1 change: 1 addition & 0 deletions litellm/proxy/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -4155,6 +4155,7 @@ class SpendLogsMetadata(TypedDict):
litellm_gateway_injected_cache: ReadOnly[str | None]
router_metadata: ReadOnly[SpendLogsRouterMetadata | None] # None = deployment not flagged internal_router_model
azure_spillover: ReadOnly[AzureSpillover | None] # None = Azure did not report spillover
used_client_oauth_token: ReadOnly[bool | None] # None = row written before the flag existed


class SpendLogsPayload(TypedDict):
Expand Down
18 changes: 17 additions & 1 deletion litellm/proxy/hooks/proxy_track_cost_callback.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
_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
Expand All @@ -29,7 +30,7 @@
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_counter_batch import post_call_counter_keys, spend_counter_batch_scope
from litellm.proxy.spend_tracking.spend_event import (
ObjectMapping,
Expand Down Expand Up @@ -86,6 +87,19 @@
)


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


def _proxy_spend_writer() -> DBSpendUpdateWriter:
from litellm.proxy.proxy_server import proxy_logging_obj

Expand Down Expand Up @@ -192,6 +206,8 @@ async def async_post_call_failure_hook(
metadata=_metadata, original_exception=original_exception
)

_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)

Expand Down
21 changes: 14 additions & 7 deletions litellm/proxy/litellm_pre_call_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
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
Expand Down Expand Up @@ -45,6 +45,7 @@
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,
Expand Down Expand Up @@ -648,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"
Expand Down Expand Up @@ -2187,7 +2191,9 @@ async def add_litellm_data_to_request(
data["api_version"] = dynamic_api_version

## Forward any LLM API Provider specific headers in extra_headers
add_provider_specific_headers_to_request(data=data, headers=_headers)
data[_metadata_variable_name]["used_client_oauth_token"] = add_provider_specific_headers_to_request(
data=data, headers=_headers
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.

## Cache Controls
cache_control_header: Final = _headers.get("Cache-Control", None)
Expand Down Expand Up @@ -3479,13 +3485,13 @@ async def add_guardrails_from_policy_engine(
LlmProviders.VERTEX_AI.value,
)
)
_ANTHROPIC_OAUTH_CREDENTIAL_PROVIDERS: Final = LlmProviders.ANTHROPIC.value
_ANTHROPIC_OAUTH_CREDENTIAL_PROVIDERS: Final = ",".join(sorted(ANTHROPIC_OAUTH_FORWARD_PROVIDERS))
Comment thread
greptile-apps[bot] marked this conversation as resolved.


def add_provider_specific_headers_to_request(
data: dict,
headers: dict,
):
) -> bool:
from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key

anthropic_api_headers: Final = {header: headers[header] for header in ANTHROPIC_API_HEADERS if header in headers}
Expand All @@ -3506,6 +3512,7 @@ def add_provider_specific_headers_to_request(

if scoped_headers:
data["provider_specific_header"] = scoped_headers[0] if len(scoped_headers) == 1 else scoped_headers
return bool(anthropic_oauth_credential_headers)


def _add_otel_traceparent_to_data(data: dict, request: Request):
Expand Down
13 changes: 13 additions & 0 deletions litellm/proxy/spend_tracking/spend_management_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -2520,6 +2520,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: 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",
Expand Down Expand Up @@ -2929,6 +2938,10 @@ def parse_date(date_str: str) -> datetime:
sql_conditions.append(f"metadata->'error_information'->>'error_message' LIKE ${p}")
sql_params.append(f"%{error_message}%")
p += 1
if used_client_oauth_token is not None:
sql_conditions.append(f"metadata->>'used_client_oauth_token' = ${p}")
sql_params.append(json.dumps(used_client_oauth_token))
p += 1

if status_filter is not None and group_by_session is True and not is_search_lookup:
session_filter_conditions: Final = " AND ".join(sql_conditions) or "TRUE"
Expand Down
9 changes: 9 additions & 0 deletions litellm/proxy/spend_tracking/spend_tracking_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
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
Expand All @@ -45,6 +46,7 @@
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.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
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
Expand Down Expand Up @@ -155,6 +157,7 @@ def _get_router_metadata_for_spend_log(
"autorouter_savings",
"autorouter_savings_estimate",
"autorouter_baseline_observation",
"used_client_oauth_token",
)
)

Expand All @@ -179,6 +182,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(
Expand Down Expand Up @@ -223,6 +227,7 @@ def _get_spend_logs_metadata(
litellm_call_id=litellm_call_id,
router_metadata=router_metadata,
azure_spillover=azure_spillover,
used_client_oauth_token=used_client_oauth_token,
)
verbose_proxy_logger.debug(
"getting payload for SpendLogs, available keys in metadata: " + str(list(metadata.keys()))
Expand All @@ -238,6 +243,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")
Expand Down Expand Up @@ -715,6 +721,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(
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")
if isinstance(kwargs.get("response_headers"), Mapping)
Expand Down
1 change: 1 addition & 0 deletions litellm/types/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3186,6 +3186,7 @@ class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata):
cold_storage_object_key: str | None # S3/GCS object key for cold storage retrieval
team_alias: str | None
team_id: str | None
used_client_oauth_token: ReadOnly[bool | None]


class AzureSpillover(TypedDict):
Expand Down
Loading
Loading