mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(proxy): resolve used_client_oauth_token against the provider the call was sent to
This commit is contained in:
parent
d688293c81
commit
05b35191d3
8 changed files with 114 additions and 16 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue