mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
* feat(anthropic): workload identity federation and pluggable identity sources Backend half of #38818 (internal copy of the fork PR #38013), rebuilt as one commit on top of litellm_internal_staging without the dashboard changes. Deployments on anthropic/ without a static api_key can exchange an OIDC workload assertion for a short-lived sk-ant-oat01 token through a shared RFC 7523 JWT-bearer engine. The assertion comes from a mounted token file, an env token, a LiteLLM-signed issuer, or Keycloak, chosen per deployment, per named credential, or through ANTHROPIC_IDENTITY_SOURCE. The federation fields are server-owned: refused inline in request bodies and on POST /model/new, proxy-admin only on credentials, and the token exchange is pinned to api.anthropic.com unless LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS adds a host. GET /credentials/{name}/jwks exports the public key set of a LiteLLM-signed credential for the Claude Console. The OpenAI federation trio from #39613 rides along on the backend side with the same server-owned handling. Fixes #28607 Resolves LIT-6107 Co-authored-by: derhornspieler <15236687+derhornspieler@users.noreply.github.com> * fix(anthropic): let batch-result downloads mint from deployment params and accept host:port allowlist entries The files handler enabled workload identity on batch-result downloads but never received the deployment's litellm_params, so a deployment authenticating through a named credential could only mint from process-wide env vars. It now threads litellm_params through to the auth header the way the batch retrieve path already does. LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS entries written as host:port were read by urlsplit as a scheme, so the allowlist kept the raw entry while the exchange compared bare hostnames and refused the gateway. Entries are now parsed as network locations whether or not they carry a scheme. * fix(types): move the WIF kwargs key sets to a leaf module so the kwargs funnel imports without a cycle * test(anthropic): pin case-insensitive matching of WIF exchange-host allowlist entries * fix(anthropic): end workload identity federation errors without a period so the router suffix reads cleanly * fix(proxy): decrypt stored litellm_params before the WIF write gate * fix(proxy): hide WIF secret references from /health output * fix(proxy): keep the proxy error shape on credential endpoint refusals * fix(proxy): hide identity token file paths from /health output * fix(anthropic): rename the federation workspace param so Bedrock's anthropic_workspace_id keeps working The Bedrock Claude Platform route already reads anthropic_workspace_id from optional_params, so banning that spelling as a server-owned federation parameter broke a pre-existing client capability. The federation field is now anthropic_federation_workspace_id (env ANTHROPIC_FEDERATION_WORKSPACE_ID), which restores the base branch's behavior for Bedrock callers, drops the Bedrock-specific hint from the refusal message, and deletes the unconditional ban constant that no longer had a reader * fix(auth): share one exchanged token across workers reading the same assertion Anthropic accepts each identity assertion exactly once, so two uvicorn workers reading the same token file both minting from it means the second exchange is denied with jti_reused. Minted tokens now land in a per-user 0700 cache directory guarded by a file lock, so workers on the same host reuse one exchange until the token expires or the assertion rotates. A 401 is only retried when the re-read assertion actually differs, and the denial hint explains jti_reused. LITELLM_TOKEN_EXCHANGE_CACHE_DIR moves the cache and an empty value disables it * fix: keep anthropic federation from being shadowed or leaked An empty or whitespace-only ANTHROPIC_API_KEY counted as set, so a federated deployment sent an empty x-api-key on every call instead of minting a token. Blank values now read as unset, and a real static key on a federated deployment logs once that it outranks federation and nothing is being federated. The exchange-host allowlist matched hostnames only, so a second process on another port of an allowed host was trusted with the workload's identity token. An entry that names a port now trusts that port alone, while a bare host still trusts every port. The shared token store exists so the workers reading one projected token file do not each spend its single-use jti. A source that mints its own assertion per exchange shares nothing with another worker, so it no longer writes a live token to disk for a lookup that can never hit. * fix: unlink a staged token file a failed write leaves behind The 401 denial hint now also says federation ignores ANTHROPIC_WORKSPACE_ID, which the Bedrock Claude platform provider already reads. * refactor: move anthropic jwks derivation behind a provider-owned tagged union * fix: unlink the staged token file when its write fails at close A buffered write only reaches the disk when the handle closes, so a full disk surfaces at close and left the staging file behind holding a usable token. * fix(anthropic): close the staging descriptor before writing the shared token file * fix(wif): judge federation writes by what they set, not what is stored The admin gate read the stored deployment, so a team admin lost edit, delete and Test Connection on any deployment carrying federation params. It now returns early unless the submitted fields touch the federation surface, and a Test Connection probe that points the deployment at its own api_base is still refused, with the 403 no longer wrapped into a 500 The rest of the same review pass: POST /model/new refuses only a blocking value of `blocked`, so a client that always sends `blocked: false` is not turned away; a request body can no longer pick which federated identity to mint as by naming a stored credential; an advisory refresh the executor refuses disarms the entry instead of wedging the identity until the follower timeout; the static-key shadow warning resolves its env fallback inside the cache instead of once per request; credential writes drop nulls before storing them; the token exchange validates the endpoint URL before reading an assertion and keeps refusing redirects across a client heal; /health hides every server-owned federation field from non-admins; and the async create_file and create_batch paths say which setting is missing when the provider resolves no URL * fix(proxy): let a deployment write name a federated credential reject_federated_credential_reference runs from is_request_body_safe, which pre_db_read_auth_checks calls on every route, so it also fired on POST /model/new, /model/update, /model/{id}/update and /health/test_connection. A proxy admin could no longer attach a federated credential to a deployment over the API or the Admin UI, leaving a static config.yaml entry as the only way to configure the feature the rejection told the caller to go configure, and _reject_non_admin_wif_write never got to make the call it exists to make. is_request_body_safe now takes the route and skips only the credential-reference check on the routes that reach can_user_make_model_call. Federation fields typed inline into a body stay refused everywhere, and a call naming a federated credential still cannot pick the identity it mints as. * refactor(proxy): derive health display policy from the federation key sets The health check module hand-copied the five workload identity fields whose value is a credential, so a shared proxy surface named provider-specific parameters and a newly added secret-bearing field would have gone on being displayed until someone remembered both places WIF_SECRET_BEARING_KEYS now sits beside the key sets it splits out of, types/utils derives secret_bearing_wif_litellm_params from it, and the health layer splats that tuple the same way it already splats the admin-only one * fix(anthropic_wif): treat blank identity-source fields as unset * test(proxy): classify the federation params in the credential slot registry main's registry test (#43298) now fails the build for any credential-named deployment param without a classification. The five federation fields that carry a token, a token file path, or a signing or client secret reference are Unplanted, matching WIF_SECRET_BEARING_KEYS; the four remaining Keycloak settings name a URL, a client id, an auth method, or a scope and are NotSecret * fix(anthropic_wif): declare federation params as owned connection leaves and chart their metrics Register the 18 Anthropic and 3 OpenAI federation params as frozen ConnectionSettings leaves so the owned-kwarg registry, the kwargs funnel and the request-body ban list read one declaration. Pass the deployment api_base through to the count-tokens handler instead of a pre-suffixed URL, which doubled the /count_tokens path on main's prompt-cache predictor. Add the five litellm_anthropic_wif_* families to the all-metrics Grafana dashboard. * fix(credentials): gate PATCH on WIF fields resolved from model_id The credential PATCH handler checked server-owned workload identity federation fields only on the values the caller sent, while a body that named a deployment through model_id had its credential values resolved after that check. A non-admin could therefore copy a federated deployment's WIF fields onto an ordinary credential. Resolve the incoming values first and run the non-admin gate on them, matching the POST path * fix(anthropic): count tokens with ANTHROPIC_AUTH_TOKEN through the shared auth header Count-tokens walked its own credential ladder: a static key, else skip minting when ANTHROPIC_AUTH_TOKEN is set, else mint a federated token. With only the auth token set it forwarded nothing and the proxy silently fell back to its local tokenizer while chat on the same deployment authenticated with that token. The handler now takes the auth header that AnthropicModelInfo.aget_auth_header resolves, the same ladder chat, files, batches and skills use, and merges the oauth beta a minted or consumer token carries with the token-counting beta --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: derhornspieler <15236687+derhornspieler@users.noreply.github.com> Co-authored-by: mateo-berri <happymvw@gmail.com>
4040 lines
157 KiB
Python
4040 lines
157 KiB
Python
"""
|
|
Unit tests for auth_utils functions related to rate limiting and customer ID extraction.
|
|
"""
|
|
|
|
import base64
|
|
import logging
|
|
from typing import Optional
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
from fastapi import HTTPException, Request
|
|
|
|
import litellm
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
from litellm.proxy.auth.auth_utils import (
|
|
_get_customer_id_from_standard_headers,
|
|
abbreviate_api_key,
|
|
check_complete_credentials,
|
|
custom_auth_common_checks_warning,
|
|
log_once_if_budget_reservation_disabled,
|
|
warn_once_if_custom_auth_skips_common_checks,
|
|
get_end_user_id_from_request_body,
|
|
get_key_mcp_rpm_limit,
|
|
get_key_model_rpm_limit,
|
|
get_key_model_tpm_limit,
|
|
get_key_own_model_rate_limit,
|
|
get_key_tag_rpm_limit,
|
|
get_model_from_request,
|
|
get_project_model_rpm_limit,
|
|
get_project_model_tpm_limit,
|
|
get_request_route_template,
|
|
is_request_body_safe,
|
|
)
|
|
from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS, OPENAI_WIF_KWARGS_KEYS
|
|
|
|
|
|
@pytest.mark.parametrize("param", sorted(ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS))
|
|
def test_every_wif_kwarg_key_is_refused_from_a_request_body(param: str):
|
|
"""Every key the kwargs funnel carries into litellm_params selects a server-side secret or the
|
|
scope a token is minted for, so each one must be refused from a request body even with the
|
|
proxy-wide client-credential opt-in; a key added to the funnel without joining the ban shows up
|
|
here as a body the proxy accepted."""
|
|
with pytest.raises(ValueError, match="server-owned workload identity federation parameter"):
|
|
is_request_body_safe(
|
|
request_body={"model": "claude-sonnet-5", param: "attacker-chosen"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="claude-sonnet-5",
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"body",
|
|
[
|
|
{"model": "claude-sonnet-5", "litellm_credential_name": "admin-wif"},
|
|
{"model": "claude-sonnet-5", "litellm_params": {"litellm_credential_name": "admin-wif"}},
|
|
],
|
|
ids=["top_level", "nested_litellm_params"],
|
|
)
|
|
def test_a_request_body_cannot_pick_a_federated_identity_by_credential_name(monkeypatch, body: dict):
|
|
"""Naming a federated credential moves the token exchange onto that credential's federation rule
|
|
and organization just as sending the fields inline does, so the ban on the inline form has to
|
|
cover the reference too."""
|
|
from litellm.types.utils import CredentialItem
|
|
|
|
monkeypatch.setattr(
|
|
litellm,
|
|
"credential_list",
|
|
[
|
|
CredentialItem(
|
|
credential_name="admin-wif",
|
|
credential_values={
|
|
"anthropic_federation_rule_id": "fdrl_admin",
|
|
"anthropic_organization_id": "org-admin",
|
|
},
|
|
credential_info={"custom_llm_provider": "anthropic"},
|
|
)
|
|
],
|
|
)
|
|
|
|
with pytest.raises(ValueError, match="names a credential configured for workload identity federation"):
|
|
is_request_body_safe(
|
|
request_body=body,
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="claude-sonnet-5",
|
|
)
|
|
|
|
|
|
def test_a_request_body_may_still_name_a_credential_that_does_not_federate(monkeypatch):
|
|
"""Only federation makes a credential a deployment decision. An ordinary named credential stays
|
|
usable from a request body, so the ban must read what the credential holds, not its presence."""
|
|
from litellm.types.utils import CredentialItem
|
|
|
|
monkeypatch.setattr(
|
|
litellm,
|
|
"credential_list",
|
|
[
|
|
CredentialItem(
|
|
credential_name="plain-key",
|
|
credential_values={"api_key": "sk-plain"},
|
|
credential_info={"custom_llm_provider": "anthropic"},
|
|
)
|
|
],
|
|
)
|
|
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "claude-sonnet-5", "litellm_credential_name": "plain-key"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="claude-sonnet-5",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def federated_credential(monkeypatch):
|
|
"""A stored credential that federates, so a body naming it is the reference the ban targets."""
|
|
from litellm.types.utils import CredentialItem
|
|
|
|
monkeypatch.setattr(
|
|
litellm,
|
|
"credential_list",
|
|
[
|
|
CredentialItem(
|
|
credential_name="admin-wif",
|
|
credential_values={
|
|
"anthropic_federation_rule_id": "fdrl_admin",
|
|
"anthropic_organization_id": "org-admin",
|
|
},
|
|
credential_info={"custom_llm_provider": "anthropic"},
|
|
)
|
|
],
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"route",
|
|
[
|
|
"/model/new",
|
|
"/model/update",
|
|
"/model/delete",
|
|
"/model/f38d7ce5-7966-42f2-bd06-67ea74aeb76b/update",
|
|
"/health/test_connection",
|
|
],
|
|
)
|
|
@pytest.mark.parametrize(
|
|
"body",
|
|
[
|
|
{"model": "claude-sonnet-5", "litellm_credential_name": "admin-wif"},
|
|
{"model": "claude-sonnet-5", "litellm_params": {"litellm_credential_name": "admin-wif"}},
|
|
],
|
|
ids=["top_level", "nested_litellm_params"],
|
|
)
|
|
def test_configuring_a_deployment_may_name_a_federated_credential(federated_credential, route: str, body: dict):
|
|
"""Attaching a federated credential to a deployment is the decision the ban tells the caller to
|
|
make, and ModelManagementAuthChecks._reject_non_admin_wif_write is what judges it: it lets a
|
|
proxy admin through and refuses everyone else with a 403. Refusing the name here first would
|
|
leave no API or Admin UI path to configure federation at all."""
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body=body,
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="claude-sonnet-5",
|
|
route=route,
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"route",
|
|
[
|
|
None,
|
|
"/v1/chat/completions",
|
|
"/v1/messages",
|
|
"/model/info",
|
|
"/model/f38d7ce5-7966-42f2-bd06-67ea74aeb76b/update/extra",
|
|
],
|
|
)
|
|
def test_a_call_still_cannot_pick_a_federated_identity_by_credential_name(federated_credential, route: str | None):
|
|
"""The exemption covers the deployment-management routes and nothing that shares their prefix,
|
|
so a call still cannot move its token exchange onto a federated credential by naming it."""
|
|
with pytest.raises(ValueError, match="names a credential configured for workload identity federation"):
|
|
is_request_body_safe(
|
|
request_body={"model": "claude-sonnet-5", "litellm_credential_name": "admin-wif"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="claude-sonnet-5",
|
|
route=route,
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize("route", ["/model/new", "/model/f38d7ce5-7966-42f2-bd06-67ea74aeb76b/update"])
|
|
def test_configuring_a_deployment_still_cannot_carry_federation_fields_inline(route: str):
|
|
"""Only the credential reference is exempt. Federation fields typed straight into a body stay
|
|
refused everywhere, since a stored credential is the surface an admin has to go through."""
|
|
with pytest.raises(ValueError, match="server-owned workload identity federation parameter"):
|
|
is_request_body_safe(
|
|
request_body={"model": "claude-sonnet-5", "anthropic_federation_rule_id": "fdrl_attacker"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="claude-sonnet-5",
|
|
route=route,
|
|
)
|
|
|
|
|
|
class TestCustomAuthCommonChecksWarning:
|
|
"""custom_auth_common_checks_warning only warns when custom auth is configured
|
|
and the common-checks opt-in is off, since that is the only state where
|
|
project/team enforcement silently does nothing."""
|
|
|
|
def test_warns_when_custom_auth_configured_and_checks_off(self):
|
|
warning = custom_auth_common_checks_warning(
|
|
custom_auth_configured=True,
|
|
run_common_checks=False,
|
|
)
|
|
assert warning is not None
|
|
assert "custom_auth_run_common_checks: true" in warning
|
|
assert "https://docs.litellm.ai/docs/proxy/custom_auth" in warning
|
|
|
|
def test_no_warning_when_common_checks_enabled(self):
|
|
assert (
|
|
custom_auth_common_checks_warning(
|
|
custom_auth_configured=True,
|
|
run_common_checks=True,
|
|
)
|
|
is None
|
|
)
|
|
|
|
def test_no_warning_when_custom_auth_not_configured(self):
|
|
assert (
|
|
custom_auth_common_checks_warning(
|
|
custom_auth_configured=False,
|
|
run_common_checks=False,
|
|
)
|
|
is None
|
|
)
|
|
assert (
|
|
custom_auth_common_checks_warning(
|
|
custom_auth_configured=False,
|
|
run_common_checks=True,
|
|
)
|
|
is None
|
|
)
|
|
|
|
|
|
class TestWarnOnceIfCustomAuthSkipsCommonChecks:
|
|
"""The startup warning must fire at most once per process, since load_config
|
|
re-runs on hot-reload / config refresh and would otherwise spam the log."""
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _reset_sentinel(self, monkeypatch):
|
|
monkeypatch.setattr(
|
|
"litellm.proxy.auth.auth_utils._custom_auth_common_checks_warning_emitted",
|
|
False,
|
|
)
|
|
|
|
def test_warns_only_once_across_repeated_calls(self):
|
|
logger = MagicMock()
|
|
for _ in range(3):
|
|
warn_once_if_custom_auth_skips_common_checks(
|
|
custom_auth_configured=True,
|
|
run_common_checks=False,
|
|
logger=logger,
|
|
)
|
|
assert logger.warning.call_count == 1
|
|
assert "custom_auth_run_common_checks" in logger.warning.call_args[0][0]
|
|
|
|
def test_does_not_warn_when_common_checks_enabled(self):
|
|
logger = MagicMock()
|
|
warn_once_if_custom_auth_skips_common_checks(
|
|
custom_auth_configured=True,
|
|
run_common_checks=True,
|
|
logger=logger,
|
|
)
|
|
assert logger.warning.call_count == 0
|
|
|
|
|
|
class TestLogOnceIfBudgetReservationDisabled:
|
|
@pytest.fixture(autouse=True)
|
|
def _reset_sentinel(self, monkeypatch):
|
|
monkeypatch.setattr(
|
|
"litellm.constants.budget_reservation_disabled_info_emitted",
|
|
False,
|
|
)
|
|
|
|
def test_logs_info_only_once_when_enabled(self, caplog):
|
|
with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"):
|
|
log_once_if_budget_reservation_disabled(disabled=False)
|
|
assert not any(
|
|
"disable_budget_reservation is enabled" in record.message
|
|
for record in caplog.records
|
|
)
|
|
for _ in range(3):
|
|
log_once_if_budget_reservation_disabled(disabled=True)
|
|
|
|
records = [
|
|
record
|
|
for record in caplog.records
|
|
if "disable_budget_reservation is enabled" in record.message
|
|
]
|
|
assert len(records) == 1
|
|
assert records[0].levelno == logging.INFO
|
|
|
|
def test_logs_to_injected_logger_only_once(self):
|
|
logger = MagicMock()
|
|
log_once_if_budget_reservation_disabled(disabled=False, logger=logger)
|
|
for _ in range(3):
|
|
log_once_if_budget_reservation_disabled(disabled=True, logger=logger)
|
|
assert logger.info.call_count == 1
|
|
assert "disable_budget_reservation is enabled" in logger.info.call_args[0][0]
|
|
|
|
|
|
class TestGetKeyModelRpmLimit:
|
|
"""Tests for get_key_model_rpm_limit function."""
|
|
|
|
def test_own_limit_excludes_team_metadata(self):
|
|
"""A team-only limit is inherited, not owned: the key resolves it but does not override it."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"some_other_key": "value"},
|
|
team_metadata={"model_rpm_limit": {"gpt-4": 50}, "model_tpm_limit": {"gpt-4": 500}},
|
|
)
|
|
assert get_key_model_rpm_limit(user_api_key_dict) == {"gpt-4": 50}
|
|
assert get_key_own_model_rate_limit(user_api_key_dict, "model_rpm_limit") is None
|
|
assert get_key_own_model_rate_limit(user_api_key_dict, "model_tpm_limit") is None
|
|
|
|
def test_own_limit_resolves_metadata_then_model_max_budget(self):
|
|
from_metadata = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"model_rpm_limit": {"gpt-4": 100}},
|
|
model_max_budget={"gpt-4": {"rpm_limit": 10, "tpm_limit": 1000}},
|
|
team_metadata={"model_rpm_limit": {"gpt-4": 50}},
|
|
)
|
|
assert get_key_own_model_rate_limit(from_metadata, "model_rpm_limit") == {"gpt-4": 100}
|
|
assert get_key_own_model_rate_limit(from_metadata, "model_tpm_limit") == {"gpt-4": 1000}
|
|
|
|
from_budget = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
model_max_budget={"gpt-4": {"rpm_limit": 10}, "gpt-3.5-turbo": {"tpm_limit": 1000}},
|
|
team_metadata={"model_rpm_limit": {"gpt-4": 50}},
|
|
)
|
|
assert get_key_own_model_rate_limit(from_budget, "model_rpm_limit") == {"gpt-4": 10}
|
|
assert get_key_own_model_rate_limit(from_budget, "model_tpm_limit") == {"gpt-3.5-turbo": 1000}
|
|
|
|
def test_returns_key_metadata_when_present(self):
|
|
"""Key metadata takes priority over team metadata."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"model_rpm_limit": {"gpt-4": 100}},
|
|
team_metadata={"model_rpm_limit": {"gpt-4": 50}},
|
|
)
|
|
result = get_key_model_rpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 100}
|
|
|
|
def test_falls_back_to_team_metadata_when_key_has_other_metadata(self):
|
|
"""Should fall back to team metadata when key metadata exists but has no model_rpm_limit."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"some_other_key": "value"}, # Has metadata, but not model_rpm_limit
|
|
team_metadata={"model_rpm_limit": {"gpt-4": 50}},
|
|
)
|
|
result = get_key_model_rpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 50}
|
|
|
|
def test_extracts_from_model_max_budget(self):
|
|
"""Should extract rpm_limit from model_max_budget when metadata is empty."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
model_max_budget={
|
|
"gpt-4": {"rpm_limit": 100, "tpm_limit": 1000},
|
|
"gpt-3.5-turbo": {"rpm_limit": 200},
|
|
},
|
|
)
|
|
result = get_key_model_rpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 100, "gpt-3.5-turbo": 200}
|
|
|
|
def test_skips_models_without_rpm_limit(self):
|
|
"""Should skip models that don't have rpm_limit in model_max_budget."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
model_max_budget={
|
|
"gpt-4": {"rpm_limit": 100},
|
|
"gpt-3.5-turbo": {"tpm_limit": 1000}, # No rpm_limit
|
|
},
|
|
)
|
|
result = get_key_model_rpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 100}
|
|
|
|
def test_returns_none_when_no_limits_configured(self):
|
|
"""Should return None when no rate limits are configured."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
result = get_key_model_rpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
def test_team_metadata_empty_rpm_dict_falls_through_to_deployment_default(self):
|
|
"""Explicitly empty team model_rpm_limit ({}) should be returned as-is, not fallen through."""
|
|
# An empty dict is a valid team limit map (no per-model limits configured).
|
|
# It should be returned directly rather than falling through to deployment defaults,
|
|
# so a team with an empty map is treated as unconstrained at the team level.
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
team_metadata={"model_rpm_limit": {}},
|
|
)
|
|
result = get_key_model_rpm_limit(user_api_key_dict)
|
|
assert result == {}
|
|
|
|
|
|
class TestGetKeyMcpRpmLimit:
|
|
def test_empty_dict_limits_are_returned(self):
|
|
key_override = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"mcp_rpm_limit": {}},
|
|
team_metadata={"mcp_rpm_limit": {"github": 50}},
|
|
)
|
|
assert get_key_mcp_rpm_limit(key_override) == {}
|
|
|
|
team_empty = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
team_metadata={"mcp_rpm_limit": {}},
|
|
)
|
|
assert get_key_mcp_rpm_limit(team_empty) == {}
|
|
|
|
|
|
class TestGetKeyModelTpmLimit:
|
|
"""Tests for get_key_model_tpm_limit function."""
|
|
|
|
def test_returns_key_metadata_when_present(self):
|
|
"""Key metadata takes priority over team metadata."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"model_tpm_limit": {"gpt-4": 10000}},
|
|
team_metadata={"model_tpm_limit": {"gpt-4": 5000}},
|
|
)
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 10000}
|
|
|
|
def test_falls_back_to_team_metadata_when_key_has_other_metadata(self):
|
|
"""Should fall back to team metadata when key metadata exists but has no model_tpm_limit."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"some_other_key": "value"}, # Has metadata, but not model_tpm_limit
|
|
team_metadata={"model_tpm_limit": {"gpt-4": 5000}},
|
|
)
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 5000}
|
|
|
|
def test_extracts_from_model_max_budget(self):
|
|
"""Should extract tpm_limit from model_max_budget when metadata is empty."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
model_max_budget={
|
|
"gpt-4": {"tpm_limit": 10000, "rpm_limit": 100},
|
|
"gpt-3.5-turbo": {"tpm_limit": 20000},
|
|
},
|
|
)
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 10000, "gpt-3.5-turbo": 20000}
|
|
|
|
def test_skips_models_without_tpm_limit(self):
|
|
"""Should skip models that don't have tpm_limit in model_max_budget."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
model_max_budget={
|
|
"gpt-4": {"tpm_limit": 10000},
|
|
"gpt-3.5-turbo": {"rpm_limit": 100}, # No tpm_limit
|
|
},
|
|
)
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 10000}
|
|
|
|
def test_returns_none_when_no_limits_configured(self):
|
|
"""Should return None when no rate limits are configured."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
def test_model_max_budget_priority_over_team(self):
|
|
"""model_max_budget should take priority over team_metadata."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
model_max_budget={"gpt-4": {"tpm_limit": 10000}},
|
|
team_metadata={"model_tpm_limit": {"gpt-4": 5000}},
|
|
)
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 10000}
|
|
|
|
def test_team_metadata_empty_tpm_dict_falls_through_to_deployment_default(self):
|
|
"""Explicitly empty team model_tpm_limit ({}) should be returned as-is, not fallen through."""
|
|
# An empty dict is a valid team limit map (no per-model limits configured).
|
|
# It should be returned directly rather than falling through to deployment defaults,
|
|
# so a team with an empty map is treated as unconstrained at the team level.
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
team_metadata={"model_tpm_limit": {}},
|
|
)
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result == {}
|
|
|
|
def test_skips_deployments_with_malformed_limit_value(self):
|
|
"""Deployments with non-integer-parseable limit values are skipped without raising."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
{
|
|
"model_name": "model1",
|
|
"litellm_params": {"default_api_key_tpm_limit": "not-a-number"},
|
|
},
|
|
_make_deployment_dict("model1", tpm=500),
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
|
|
# The malformed deployment is skipped; the valid one provides 500
|
|
assert result == {"model1": 500}
|
|
|
|
|
|
class TestGetCustomerIdFromStandardHeaders:
|
|
"""Tests for _get_customer_id_from_standard_headers helper function."""
|
|
|
|
def test_should_return_customer_id_from_x_litellm_customer_id_header(self):
|
|
"""Should extract customer ID from x-litellm-customer-id header."""
|
|
headers = {"x-litellm-customer-id": "customer-123"}
|
|
result = _get_customer_id_from_standard_headers(request_headers=headers)
|
|
assert result == "customer-123"
|
|
|
|
def test_should_return_customer_id_from_x_litellm_end_user_id_header(self):
|
|
"""Should extract customer ID from x-litellm-end-user-id header."""
|
|
headers = {"x-litellm-end-user-id": "end-user-456"}
|
|
result = _get_customer_id_from_standard_headers(request_headers=headers)
|
|
assert result == "end-user-456"
|
|
|
|
def test_should_return_none_when_headers_is_none(self):
|
|
"""Should return None when headers is None."""
|
|
result = _get_customer_id_from_standard_headers(request_headers=None)
|
|
assert result is None
|
|
|
|
def test_should_return_none_when_no_standard_headers_present(self):
|
|
"""Should return None when no standard customer ID headers are present."""
|
|
headers = {"x-other-header": "some-value"}
|
|
result = _get_customer_id_from_standard_headers(request_headers=headers)
|
|
assert result is None
|
|
|
|
|
|
class TestGetEndUserIdFromRequestBodyWithStandardHeaders:
|
|
"""Tests for get_end_user_id_from_request_body with standard customer ID headers."""
|
|
|
|
def test_should_prioritize_standard_header_over_body_user(self):
|
|
"""Standard customer ID header should take precedence over body user field."""
|
|
headers = {"x-litellm-customer-id": "header-customer"}
|
|
request_body = {"user": "body-user"}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers)
|
|
assert result == "header-customer"
|
|
|
|
def test_should_fall_back_to_body_when_no_standard_header(self):
|
|
"""Should fall back to body user when no standard headers are present."""
|
|
headers = {"x-other-header": "value"}
|
|
request_body = {"user": "body-user"}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers)
|
|
assert result == "body-user"
|
|
|
|
|
|
def _request_dispatched_to(endpoint) -> Request:
|
|
"""Build a minimal Request whose FastAPI-resolved endpoint is ``endpoint``,
|
|
mirroring what Starlette sets in ``scope`` once routing has matched."""
|
|
return Request(scope={"type": "http", "headers": [], "endpoint": endpoint})
|
|
|
|
|
|
def _pass_through_endpoint():
|
|
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
|
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
|
|
)
|
|
|
|
def endpoint(): # stand-in for create_pass_through_route's handler
|
|
...
|
|
|
|
setattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
|
|
return endpoint
|
|
|
|
|
|
def test_get_model_from_request_skips_pass_through_dispatched_request():
|
|
"""When FastAPI dispatched the request to a user-defined pass-through handler,
|
|
the body `model` names an upstream model and must not be treated as a LiteLLM
|
|
model for allowlist/budget enforcement."""
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "upstream-special-model"},
|
|
route="/my-custom-endpoint",
|
|
request=_request_dispatched_to(_pass_through_endpoint()),
|
|
)
|
|
is None
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_enforces_when_builtin_handler_dispatched():
|
|
"""A custom pass-through path that collides with a built-in route resolves to the
|
|
built-in handler (no marker), so the body `model` must still be extracted and
|
|
enforced. Same request path as above, but dispatched to a non-pass-through
|
|
endpoint: the model must NOT be suppressed."""
|
|
|
|
def builtin_chat_completions(): ...
|
|
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "gpt-4o"},
|
|
route="/v1/chat/completions",
|
|
request=_request_dispatched_to(builtin_chat_completions),
|
|
)
|
|
== "gpt-4o"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_no_request_extracts_model():
|
|
"""Callers without a request object (e.g. budget reservation) still extract the
|
|
model; the pass-through suppression only applies to a dispatched pass-through
|
|
handler."""
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "gpt-4o"},
|
|
route="/v1/chat/completions",
|
|
)
|
|
== "gpt-4o"
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize("provider,model", [
|
|
("laya", "english"), ("laya", "multilingual"), ("laya", "typed-decisions"),
|
|
("bespoke", "nimble-latest"), ("bespoke", "bespokelabs/Bespoke-Nimble-9B"),
|
|
])
|
|
@pytest.mark.parametrize("suffix", ["", "/"])
|
|
def test_oss_native_model_uses_the_classifier_permission_identity(provider: str, model: str, suffix: str) -> None:
|
|
assert get_model_from_request(
|
|
request_data={"model": model}, route=f"/{provider}/v1/systemone{suffix}"
|
|
) == f"{provider}/{model}"
|
|
|
|
|
|
@pytest.mark.parametrize("provider", ["laya", "bespoke"])
|
|
@pytest.mark.parametrize("model", [None, "", "auto", "laya/english", "bespoke/nimble-latest", "unknown", ["english"], 7])
|
|
def test_oss_native_model_cannot_implicitly_select_an_unauthorized_checkpoint(provider: str, model: object) -> None:
|
|
with pytest.raises(HTTPException) as denied:
|
|
get_model_from_request(request_data={"model": model}, route=f"/{provider}/v1/systemone")
|
|
assert denied.value.status_code == 400
|
|
|
|
|
|
def test_laya_model_normalization_does_not_change_other_provider_routes() -> None:
|
|
assert get_model_from_request(request_data={"model": "jev-latest"}, route="/typesafe/v1/systemone") == "jev-latest"
|
|
assert get_model_from_request(request_data={}, route="/laya/health") is None
|
|
|
|
|
|
def _cache_prediction_router():
|
|
from litellm.router import Router
|
|
|
|
return Router(model_list=[
|
|
{
|
|
"model_name": group,
|
|
"litellm_params": {"model": "anthropic/claude-sonnet-5", "api_key": "test-provider-key"},
|
|
"model_info": {"id": deployment_id, "team_id": team_id},
|
|
}
|
|
for group, deployment_id, team_id in (
|
|
("current-group", "current-id", None), ("candidate-group", "candidate-id", None),
|
|
("own-group", "own-id", "prediction-team"), ("foreign-group", "foreign-id", "foreign-team"),
|
|
)
|
|
])
|
|
|
|
|
|
@pytest.mark.parametrize("candidate,team_id,expected", [
|
|
("candidate-id", None, ["current-group", "candidate-group"]),
|
|
("current-id", None, "current-group"),
|
|
("missing-id", None, None),
|
|
("candidate-group", None, None),
|
|
("own-id", None, None),
|
|
("own-id", "prediction-team", ["current-group", "own-group"]),
|
|
("foreign-id", "prediction-team", None),
|
|
])
|
|
def test_cache_prediction_auth_resolves_only_exact_deployment_ids(candidate, team_id, expected):
|
|
assert get_model_from_request(
|
|
request_data={
|
|
"current_deployment_id": "current-id", "candidate_deployment_id": candidate,
|
|
"request": {"model": "caller-controlled-provider-model"},
|
|
},
|
|
route="/cost/predict-cache",
|
|
llm_router=_cache_prediction_router(),
|
|
team_id=team_id,
|
|
) == expected
|
|
|
|
|
|
def _cache_prediction_auth_app(
|
|
monkeypatch, allowed_routes, user_models, metadata=None, *, team_id=None, key_models=None, team_models=None
|
|
):
|
|
import importlib
|
|
from unittest.mock import AsyncMock
|
|
|
|
from fastapi import FastAPI
|
|
|
|
import litellm.proxy.proxy_server as proxy_server
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LitellmUserRoles, ProxyException
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
|
|
from litellm.proxy.management_endpoints import prompt_cache_prediction as endpoint
|
|
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
|
|
|
|
auth = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
|
|
router = _cache_prediction_router()
|
|
allowed_models = ["current-group", "candidate-group", "own-group"]
|
|
token = UserAPIKeyAuth(
|
|
api_key="test-proxy-key-hash", user_id="prediction-user", user_role=LitellmUserRoles.INTERNAL_USER,
|
|
models=allowed_models if key_models is None else key_models, team_id=team_id,
|
|
team_models=allowed_models if team_models is None else team_models,
|
|
allowed_routes=allowed_routes, metadata=metadata or {},
|
|
)
|
|
user = LiteLLM_UserTable(
|
|
user_id=token.user_id, user_role=LitellmUserRoles.INTERNAL_USER.value, models=user_models,
|
|
)
|
|
async def authenticate(request, request_data, **_headers):
|
|
await auth._enforce_key_and_fallback_model_access(
|
|
valid_token=token, request_data=request_data, route=request.url.path, request=request,
|
|
llm_model_list=router.get_model_list(), llm_router=router,
|
|
)
|
|
return token
|
|
|
|
monkeypatch.setattr(auth, "_user_api_key_auth_builder", authenticate)
|
|
monkeypatch.setattr(auth, "get_user_object", AsyncMock(return_value=user))
|
|
team = LiteLLM_TeamTableCachedObj(team_id=team_id, models=token.team_models) if team_id else None
|
|
monkeypatch.setattr(auth, "get_team_object", AsyncMock(return_value=team))
|
|
monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=team))
|
|
monkeypatch.setattr(auth_checks, "get_team_membership", AsyncMock(return_value=None))
|
|
monkeypatch.setattr(auth, "get_global_proxy_spend", AsyncMock(return_value=0))
|
|
monkeypatch.setattr(proxy_server, "master_key", "test-master-key")
|
|
monkeypatch.setattr(proxy_server, "user_custom_auth", None)
|
|
monkeypatch.setattr(proxy_server, "general_settings", {})
|
|
monkeypatch.setattr(proxy_server, "llm_router", router)
|
|
monkeypatch.setattr(proxy_server, "llm_model_list", router.get_model_list())
|
|
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
|
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
|
|
logging = ProxyLogging(user_api_key_cache=DualCache())
|
|
logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3(
|
|
InternalUsageCache(dual_cache=DualCache())
|
|
)
|
|
monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging)
|
|
counts = AsyncMock(return_value=6_000)
|
|
monkeypatch.setattr(endpoint, "count_prompt_tokens", counts)
|
|
app = FastAPI()
|
|
app.include_router(endpoint.router)
|
|
app.add_exception_handler(ProxyException, proxy_server.openai_exception_handler)
|
|
return app, counts
|
|
|
|
|
|
def _cache_prediction_payload(candidate="candidate-id", current="current-id"):
|
|
return {
|
|
"current_deployment_id": current, "candidate_deployment_id": candidate,
|
|
"request": {"messages": [{"role": "user", "content": [{
|
|
"type": "text", "text": "Stable cached context",
|
|
"cache_control": {"type": "ephemeral"},
|
|
}]}]},
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("allowed_routes,user_models,candidate,status_code", [
|
|
(["/chat/completions"], ["current-group", "candidate-group"], "candidate-id", 403),
|
|
(["/cost/predict-cache"], ["current-group"], "candidate-id", 403),
|
|
(["/cost/*"], ["current-group", "candidate-group"], "candidate-id", 200),
|
|
(["/cost/predict-cache"], ["current-group"], "current-id", 200),
|
|
(["/cost/predict-cache"], ["current-group"], "missing-id", 404),
|
|
])
|
|
async def test_cache_prediction_authorizes_route_and_personal_models_before_provider_counts(
|
|
monkeypatch, allowed_routes, user_models, candidate, status_code
|
|
):
|
|
import httpx
|
|
|
|
app, counts = _cache_prediction_auth_app(monkeypatch, allowed_routes, user_models)
|
|
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
|
|
response = await client.post("/cost/predict-cache", json=_cache_prediction_payload(candidate))
|
|
|
|
assert response.status_code == status_code, response.text
|
|
if status_code == 200:
|
|
assert counts.await_count == (2 if candidate == "current-id" else 4)
|
|
else:
|
|
assert counts.await_count == 0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("arm", ["current_deployment_id", "candidate_deployment_id"])
|
|
@pytest.mark.parametrize("team_id,key_models,user_models,team_models", [
|
|
(None, ["*"], ["*"], None),
|
|
(None, ["current-group", "candidate-group"], ["*"], None),
|
|
(None, ["*"], ["current-group", "candidate-group"], None),
|
|
("prediction-team", ["*"], ["*"], ["current-group", "candidate-group"]),
|
|
])
|
|
async def test_cache_prediction_hides_foreign_and_missing_ids_before_model_authorization(
|
|
monkeypatch, arm, team_id, key_models, user_models, team_models
|
|
):
|
|
import httpx
|
|
|
|
app, counts = _cache_prediction_auth_app(
|
|
monkeypatch, ["/cost/predict-cache"], user_models,
|
|
team_id=team_id, key_models=key_models, team_models=team_models,
|
|
)
|
|
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
|
|
missing = await client.post("/cost/predict-cache", json={**_cache_prediction_payload(), arm: "missing-id"})
|
|
foreign = await client.post("/cost/predict-cache", json={**_cache_prediction_payload(), arm: "foreign-id"})
|
|
|
|
assert missing.status_code == foreign.status_code == 404, foreign.text
|
|
assert missing.json() == foreign.json() == {"detail": "Deployment not found"}
|
|
assert counts.await_count == 0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("arm", ["current_deployment_id", "candidate_deployment_id"])
|
|
@pytest.mark.parametrize("key_models,team_models,status_code", [
|
|
(["*"], ["*"], 200),
|
|
(["current-group", "candidate-group"], ["*"], 403),
|
|
(["*"], ["current-group", "candidate-group"], 403),
|
|
])
|
|
async def test_cache_prediction_checks_each_visible_team_deployment_model(
|
|
monkeypatch, arm, key_models, team_models, status_code
|
|
):
|
|
import httpx
|
|
|
|
app, counts = _cache_prediction_auth_app(
|
|
monkeypatch, ["/cost/predict-cache"], ["*"],
|
|
team_id="prediction-team", key_models=key_models, team_models=team_models,
|
|
)
|
|
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
|
|
response = await client.post("/cost/predict-cache", json={**_cache_prediction_payload(), arm: "own-id"})
|
|
|
|
assert response.status_code == status_code, response.text
|
|
assert counts.await_count == (4 if status_code == 200 else 0)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("arm", ["current_deployment_id", "candidate_deployment_id"])
|
|
async def test_cache_prediction_checks_each_visible_personal_deployment_model(monkeypatch, arm):
|
|
import httpx
|
|
|
|
app, counts = _cache_prediction_auth_app(monkeypatch, ["/cost/predict-cache"], ["current-group"])
|
|
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
|
|
response = await client.post(
|
|
"/cost/predict-cache", json={**_cache_prediction_payload(candidate="current-id"), arm: "candidate-id"}
|
|
)
|
|
|
|
assert response.status_code == 403, response.text
|
|
assert response.json()["error"]["type"] == "user_model_access_denied"
|
|
assert counts.await_count == 0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("header_tag,key_tags,limit,status_code,provider_calls", [
|
|
("limited", [], 1, 429, 1),
|
|
(None, ["limited"], 1, 429, 1),
|
|
("limited", ["limited"], 4, 200, 4),
|
|
("unlimited", [], 1, 200, 4),
|
|
])
|
|
async def test_cache_prediction_preserves_authenticated_header_and_key_tag_rpm(
|
|
monkeypatch, header_tag, key_tags, limit, status_code, provider_calls
|
|
):
|
|
import httpx
|
|
|
|
app, counts = _cache_prediction_auth_app(
|
|
monkeypatch, ["/cost/predict-cache"], ["current-group", "candidate-group"],
|
|
metadata={"tag_rpm_limit": {"limited": limit}, "tags": key_tags},
|
|
)
|
|
headers = {"x-litellm-tags": header_tag} if header_tag else {}
|
|
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
|
|
response = await client.post("/cost/predict-cache", json=_cache_prediction_payload(), headers=headers)
|
|
assert response.status_code == status_code, response.text
|
|
assert counts.await_count == provider_calls
|
|
if limit == 4:
|
|
exhausted = await client.post("/cost/predict-cache", json=_cache_prediction_payload(), headers=headers)
|
|
assert exhausted.status_code == 429, exhausted.text
|
|
assert counts.await_count == 4
|
|
assert all("metadata" not in call.args[2] for call in counts.await_args_list)
|
|
|
|
|
|
def test_get_model_from_request_supports_google_model_names_with_slashes():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/v1beta/models/bedrock/claude-sonnet-3.7:generateContent",
|
|
)
|
|
== "bedrock/claude-sonnet-3.7"
|
|
)
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/models/hosted_vllm/gpt-oss-20b:generateContent",
|
|
)
|
|
== "hosted_vllm/gpt-oss-20b"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_vertex_passthrough_still_works():
|
|
route = "/vertex_ai/v1/projects/p/locations/l/publishers/google/models/gemini-1.5-pro:generateContent"
|
|
assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro"
|
|
|
|
|
|
def test_get_model_from_request_openai_deployment_route_still_works():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/openai/deployments/my-azure-deployment/chat/completions",
|
|
)
|
|
== "my-azure-deployment"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_converse_passthrough():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/bedrock/model/us.anthropic.claude-sonnet-4-6/converse",
|
|
)
|
|
== "us.anthropic.claude-sonnet-4-6"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_invoke_passthrough():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke",
|
|
)
|
|
== "us.anthropic.claude-sonnet-4-6"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_v2_converse_stream_passthrough():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/bedrock/v2/model/us.anthropic.claude-sonnet-4-6/converse-stream",
|
|
)
|
|
== "us.anthropic.claude-sonnet-4-6"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_model_id_with_slashes():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/bedrock/model/aws/anthropic/model-name/invoke",
|
|
)
|
|
== "aws/anthropic/model-name"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_unparseable_endpoint_returns_none():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/bedrock/agents/some-agent-route",
|
|
)
|
|
is None
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_url_model_overrides_body_model():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "us.anthropic.claude-sonnet-4-6"},
|
|
route="/bedrock/model/us.anthropic.claude-opus-4-6-v1/converse",
|
|
)
|
|
== "us.anthropic.claude-opus-4-6-v1"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_invoke_url_model_overrides_body_model():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "us.anthropic.claude-sonnet-4-6"},
|
|
route="/bedrock/model/us.anthropic.claude-opus-4-6-v1/invoke",
|
|
)
|
|
== "us.anthropic.claude-opus-4-6-v1"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_count_tokens_uses_body_model():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "us.anthropic.claude-sonnet-4-6"},
|
|
route="/bedrock/v1/messages/count_tokens",
|
|
)
|
|
== "us.anthropic.claude-sonnet-4-6"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_uppercase_count_tokens_segment_is_not_count_tokens():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "us.anthropic.claude-haiku-4-5-20251001-v1:0"},
|
|
route="/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke/COUNT_TOKENS",
|
|
)
|
|
== "us.anthropic.claude-sonnet-4-6"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_bedrock_unparseable_endpoint_keeps_body_model():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "us.anthropic.claude-sonnet-4-6"},
|
|
route="/bedrock/agents/some-agent-route",
|
|
)
|
|
== "us.anthropic.claude-sonnet-4-6"
|
|
)
|
|
|
|
|
|
def _azure_relay_router():
|
|
from litellm.router import Router
|
|
|
|
return Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt",
|
|
"litellm_params": {"model": "azure_ai/gpt-5.4-mini", "api_base": "https://a.services.ai.azure.com", "api_key": "k"},
|
|
},
|
|
{
|
|
"model_name": "other-group",
|
|
"litellm_params": {"model": "azure/gpt-5.4", "api_base": "https://b.openai.azure.com", "api_key": "k"},
|
|
},
|
|
]
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"route, request_data, expected",
|
|
[
|
|
("/azure_ai/other-group/openai/deployments/other-group/chat/completions", {"model": "gpt"}, "other-group"),
|
|
("/azure_ai/other-group/models/chat/completions", {}, "other-group"),
|
|
("/azure/openai/deployments/gpt/chat/completions", {"model": "other-group"}, "gpt"),
|
|
("/azure/openai/deployments/gpt/chat/completions", {}, "gpt"),
|
|
("/azure/openai/deployments/my-azure-deployment/chat/completions", {"model": "gpt"}, "gpt"),
|
|
("/azure_ai/gpt", {"model": "other-group"}, "other-group"),
|
|
],
|
|
)
|
|
def test_get_model_from_request_azure_relay_routes_use_the_model_group_in_the_path(route, request_data, expected):
|
|
assert get_model_from_request(request_data=request_data, route=route, llm_router=_azure_relay_router()) == expected
|
|
|
|
|
|
def _nvidia_nim_relay_router():
|
|
from litellm.router import Router
|
|
|
|
return Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "nim-page-elements",
|
|
"litellm_params": {
|
|
"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2",
|
|
"api_base": "http://nim-a.internal:8000",
|
|
"api_key": "k",
|
|
},
|
|
},
|
|
{
|
|
"model_name": "nvidia/nemoretriever-table-structure-v1",
|
|
"litellm_params": {
|
|
"model": "nvidia_nim/nvidia/nemoretriever-table-structure-v1",
|
|
"api_base": "http://nim-b.internal:8000",
|
|
"api_key": "k",
|
|
},
|
|
},
|
|
{
|
|
"model_name": "gpt-4o",
|
|
"litellm_params": {"model": "openai/gpt-4o", "api_key": "k"},
|
|
},
|
|
{
|
|
"model_name": "detect",
|
|
"litellm_params": {
|
|
"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2",
|
|
"api_base": "http://nim-a.internal:8000",
|
|
"api_key": "k",
|
|
},
|
|
},
|
|
{
|
|
"model_name": "detect",
|
|
"litellm_params": {"model": "openai/gpt-4o", "api_key": "k"},
|
|
},
|
|
]
|
|
)
|
|
|
|
|
|
NIM_INFER_BODY = {"input": [{"type": "image_url", "url": "data:image/png;base64,AAAA"}]}
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"route, request_data, expected",
|
|
[
|
|
("/nvidia_nim/nim-page-elements/v1/infer", NIM_INFER_BODY, "nim-page-elements"),
|
|
(
|
|
"/nvidia_nim/nim-page-elements/v1/infer",
|
|
{"model": "nvidia/nemoretriever-table-structure-v1"},
|
|
"nim-page-elements",
|
|
),
|
|
(
|
|
"/nvidia_nim/nvidia/nemoretriever-table-structure-v1/v1/infer",
|
|
NIM_INFER_BODY,
|
|
"nvidia/nemoretriever-table-structure-v1",
|
|
),
|
|
("/nvidia_nim/v1/infer", NIM_INFER_BODY, None),
|
|
("/nvidia_nim/unknown-group/v1/infer", NIM_INFER_BODY, None),
|
|
("/nvidia_nim/nim-page-elements-v2/v1/infer", NIM_INFER_BODY, None),
|
|
("/nvidia_nim/gpt-4o/v1/infer", NIM_INFER_BODY, None),
|
|
("/nvidia_nim/detect/v1/infer", NIM_INFER_BODY, None),
|
|
],
|
|
)
|
|
def test_get_model_from_request_nvidia_nim_relay_routes_use_the_model_group_in_the_path(route, request_data, expected):
|
|
assert (
|
|
get_model_from_request(request_data=request_data, route=route, llm_router=_nvidia_nim_relay_router())
|
|
== expected
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_nvidia_nim_relay_without_a_router_has_no_model():
|
|
assert get_model_from_request(request_data=NIM_INFER_BODY, route="/nvidia_nim/nim-page-elements/v1/infer") is None
|
|
|
|
|
|
def test_get_model_from_request_includes_file_endpoint_header_model():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/v1/files",
|
|
request_headers={"X-LiteLLM-Model": "restricted-model"},
|
|
)
|
|
== "restricted-model"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_ignores_routing_header_on_standard_llm_routes():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"model": "allowed-model"},
|
|
route="/v1/chat/completions",
|
|
request_headers={"x-litellm-model": "restricted-model"},
|
|
)
|
|
== "allowed-model"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_authorizes_all_file_routing_model_sources():
|
|
models = get_model_from_request(
|
|
request_data={"model": "body-model"},
|
|
route="/v1/files",
|
|
request_headers={"x-litellm-model": "header-model"},
|
|
request_query_params={"target_model_names": "query-model-a,query-model-b"},
|
|
)
|
|
assert isinstance(models, list)
|
|
assert set(models) == {
|
|
"body-model",
|
|
"query-model-a",
|
|
"query-model-b",
|
|
"header-model",
|
|
}
|
|
|
|
|
|
def test_get_model_from_request_extracts_simple_encoded_file_id_model():
|
|
from litellm.proxy.openai_files_endpoints.common_utils import (
|
|
encode_file_id_with_model,
|
|
)
|
|
|
|
file_id = encode_file_id_with_model(
|
|
file_id="file-provider-id",
|
|
model="restricted-model",
|
|
)
|
|
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"file_id": file_id},
|
|
route="/v1/files/{file_id}",
|
|
)
|
|
== "restricted-model"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_extracts_unified_file_id_models():
|
|
raw_unified_file_id = (
|
|
"litellm_proxy:application/octet-stream;unified_id,test-id;"
|
|
"target_model_names,model-a,model-b;llm_output_file_id,file-provider-id"
|
|
)
|
|
encoded_unified_file_id = base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=")
|
|
|
|
assert get_model_from_request(
|
|
request_data={"file_id": encoded_unified_file_id},
|
|
route="/v1/files/{file_id}",
|
|
) == ["model-a", "model-b"]
|
|
|
|
|
|
def test_get_model_from_request_extracts_eval_completion_model():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"completion": {"model": "judge-model"}},
|
|
route="/v1/evals/{eval_id}/runs",
|
|
)
|
|
== "judge-model"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_includes_fine_tuning_target_model_query():
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={},
|
|
route="/v1/fine_tuning/jobs",
|
|
request_query_params={"target_model_names": "fine-tune-model"},
|
|
)
|
|
== "fine-tune-model"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_extracts_video_id_model():
|
|
from litellm.types.videos.utils import encode_video_id_with_provider
|
|
|
|
video_id = encode_video_id_with_provider(
|
|
video_id="video-provider-id",
|
|
provider="openai",
|
|
model_id="video-model",
|
|
)
|
|
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"video_id": video_id},
|
|
route="/v1/videos/{video_id}",
|
|
)
|
|
== "video-model"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_resolves_video_id_model_with_router():
|
|
from litellm.types.videos.utils import encode_video_id_with_provider
|
|
|
|
provider_video_id = (
|
|
"projects/test-project/locations/us-central1/publishers/google/models/"
|
|
"veo-3.1-generate-001/operations/operation-id"
|
|
)
|
|
video_id = encode_video_id_with_provider(
|
|
video_id=provider_video_id,
|
|
provider="vertex_ai",
|
|
model_id="veo-3.1-generate-001",
|
|
)
|
|
llm_router = MagicMock()
|
|
llm_router.resolve_model_name_from_model_id.return_value = "gcp/google/veo-3.1-generate-001"
|
|
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"video_id": video_id},
|
|
route="/v1/videos/{video_id}",
|
|
llm_router=llm_router,
|
|
)
|
|
== "gcp/google/veo-3.1-generate-001"
|
|
)
|
|
llm_router.resolve_model_name_from_model_id.assert_called_once_with("veo-3.1-generate-001")
|
|
|
|
|
|
_BATCH_DEPLOYMENT_ID = "8d0eaa7e6c6f54a425dfd0062cb6b0dc"
|
|
|
|
|
|
def _managed_batch_router():
|
|
from litellm.router import Router
|
|
|
|
return Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "bedrock-batch-model",
|
|
"litellm_params": {
|
|
"model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0",
|
|
},
|
|
"model_info": {"id": _BATCH_DEPLOYMENT_ID},
|
|
},
|
|
{
|
|
"model_name": "some-other-model",
|
|
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
|
|
"model_info": {"id": "a-different-deployment-id"},
|
|
},
|
|
]
|
|
)
|
|
|
|
|
|
def _encode_managed_id(decoded: str) -> str:
|
|
return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=")
|
|
|
|
|
|
_MANAGED_BATCH_ID = _encode_managed_id(f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123")
|
|
_MANAGED_BATCH_OUTPUT_FILE_ID = _encode_managed_id(
|
|
f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123;"
|
|
"llm_output_file_id:provider-file-456"
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"route, request_data",
|
|
[
|
|
("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}),
|
|
("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}),
|
|
("/v1/files/{file_id}", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}),
|
|
("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}),
|
|
],
|
|
)
|
|
def test_get_model_from_request_resolves_batch_id_deployment_to_model_name(route, request_data):
|
|
"""Regression for #32580: managed batch retrieve/cancel and managed batch output
|
|
file reads encode the deployment model_id into the resource id. The auth layer must
|
|
resolve that id back to the public model group name so model-access checks compare
|
|
against the model group, not the raw deployment id."""
|
|
assert (
|
|
get_model_from_request(
|
|
request_data=request_data,
|
|
route=route,
|
|
llm_router=_managed_batch_router(),
|
|
)
|
|
== "bedrock-batch-model"
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"route, request_data",
|
|
[
|
|
("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}),
|
|
("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}),
|
|
("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}),
|
|
],
|
|
)
|
|
async def test_managed_batch_routes_pass_team_model_access_check(route, request_data):
|
|
"""End-to-end regression for #32580: a team scoped to the batch model group got
|
|
``team_model_access_denied`` on retrieve/cancel because the deployment id, not the
|
|
model group, was authorized. Fails pre-fix with the deployment id in the message."""
|
|
from litellm.proxy._types import LiteLLM_TeamTable
|
|
from litellm.proxy.auth.auth_checks import can_team_access_model
|
|
|
|
llm_router = _managed_batch_router()
|
|
model = get_model_from_request(request_data=request_data, route=route, llm_router=llm_router)
|
|
|
|
assert (
|
|
await can_team_access_model(
|
|
model=model,
|
|
team_object=LiteLLM_TeamTable(team_id="team-batch", models=["bedrock-batch-model"]),
|
|
llm_router=llm_router,
|
|
)
|
|
is True
|
|
)
|
|
|
|
with pytest.raises(Exception, match="is not available for this API key"):
|
|
await can_team_access_model(
|
|
model=model,
|
|
team_object=LiteLLM_TeamTable(team_id="team-other", models=["some-other-model"]),
|
|
llm_router=llm_router,
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_resolves_character_id_model_with_router():
|
|
from litellm.types.videos.utils import encode_character_id_with_provider
|
|
|
|
character_id = encode_character_id_with_provider(
|
|
character_id="character-provider-id",
|
|
provider="vertex_ai",
|
|
model_id="veo-3.1-generate-001",
|
|
)
|
|
llm_router = MagicMock()
|
|
llm_router.resolve_model_name_from_model_id.return_value = "gcp/google/veo-3.1-generate-001"
|
|
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"character_id": character_id},
|
|
route="/v1/videos/characters/{character_id}",
|
|
llm_router=llm_router,
|
|
)
|
|
== "gcp/google/veo-3.1-generate-001"
|
|
)
|
|
llm_router.resolve_model_name_from_model_id.assert_called_once_with("veo-3.1-generate-001")
|
|
|
|
|
|
def test_get_model_from_request_only_runs_media_decoders_for_matching_fields():
|
|
with (
|
|
patch(
|
|
"litellm.types.videos.utils.decode_video_id_with_provider",
|
|
return_value={"model_id": "video-model"},
|
|
) as video_decoder,
|
|
patch(
|
|
"litellm.types.videos.utils.decode_character_id_with_provider",
|
|
return_value={"model_id": "character-model"},
|
|
) as character_decoder,
|
|
):
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"file_id": "file-provider-id"},
|
|
route="/v1/files/{file_id}",
|
|
)
|
|
is None
|
|
)
|
|
video_decoder.assert_not_called()
|
|
character_decoder.assert_not_called()
|
|
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"video_id": "video-provider-id"},
|
|
route="/v1/videos/{video_id}",
|
|
)
|
|
== "video-model"
|
|
)
|
|
video_decoder.assert_called_once_with("video-provider-id")
|
|
character_decoder.assert_not_called()
|
|
|
|
video_decoder.reset_mock()
|
|
character_decoder.reset_mock()
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"character_id": "character-provider-id"},
|
|
route="/v1/videos/{character_id}",
|
|
)
|
|
== "character-model"
|
|
)
|
|
video_decoder.assert_not_called()
|
|
character_decoder.assert_called_once_with("character-provider-id")
|
|
|
|
|
|
def test_get_model_from_request_handles_managed_id_decoder_failures():
|
|
with (
|
|
patch(
|
|
"litellm.proxy.openai_files_endpoints.common_utils.decode_model_from_file_id",
|
|
side_effect=Exception("decode failed"),
|
|
),
|
|
patch(
|
|
"litellm.llms.base_llm.managed_resources.utils.parse_unified_id",
|
|
side_effect=Exception("parse failed"),
|
|
),
|
|
patch(
|
|
"litellm.types.videos.utils.decode_video_id_with_provider",
|
|
side_effect=Exception("video decode failed"),
|
|
),
|
|
):
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"file_id": "not-a-managed-resource-id"},
|
|
route="/v1/files/{file_id}",
|
|
)
|
|
is None
|
|
)
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"video_id": "not-a-managed-resource-id"},
|
|
route="/v1/videos/{video_id}",
|
|
)
|
|
is None
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"route",
|
|
[
|
|
"/realtime/client_secrets",
|
|
"/v1/realtime/client_secrets",
|
|
"/openai/v1/realtime/client_secrets",
|
|
"/realtime/calls",
|
|
"/v1/realtime/calls",
|
|
"/openai/v1/realtime/calls",
|
|
],
|
|
)
|
|
def test_get_model_from_request_extracts_realtime_session_model(route):
|
|
"""The effective realtime model lives in ``session.model`` (not the
|
|
top-level ``model``). It must be surfaced so can_key_call_model() can
|
|
validate the model a restricted key is actually requesting.
|
|
|
|
Regression test for the model-access bypass on the GA Realtime WebRTC
|
|
HTTP routes (https://github.com/BerriAI/litellm/issues/29923).
|
|
"""
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"session": {"type": "realtime", "model": "gpt-realtime"}},
|
|
route=route,
|
|
)
|
|
== "gpt-realtime"
|
|
)
|
|
|
|
|
|
def test_get_model_from_request_realtime_includes_top_level_and_session_model():
|
|
"""When both top-level and session model are present, both are returned so
|
|
neither path can smuggle a disallowed model past the model-access check."""
|
|
models = get_model_from_request(
|
|
request_data={
|
|
"model": "gpt-4o-realtime-preview",
|
|
"session": {"type": "realtime", "model": "gpt-realtime"},
|
|
},
|
|
route="/v1/realtime/client_secrets",
|
|
)
|
|
assert models == ["gpt-4o-realtime-preview", "gpt-realtime"]
|
|
|
|
|
|
def test_get_model_from_request_ignores_session_model_on_non_realtime_routes():
|
|
"""A nested ``session.model`` must not leak into model resolution for
|
|
unrelated routes."""
|
|
assert (
|
|
get_model_from_request(
|
|
request_data={"session": {"type": "realtime", "model": "gpt-realtime"}},
|
|
route="/v1/chat/completions",
|
|
)
|
|
is None
|
|
)
|
|
|
|
|
|
def test_abbreviate_api_key():
|
|
assert abbreviate_api_key("sk-test-1234-abcdefgh") == "sk-...efgh"
|
|
assert abbreviate_api_key("sk-abcdefghijklm") == "sk-...jklm"
|
|
|
|
|
|
def test_abbreviate_api_key_short_key_is_fully_masked():
|
|
"""Regression test for LIT-4355: for keys shorter than the enforced minimum,
|
|
showing the last 4 characters can reveal the entire key (sk-1234 -> sk-...1234)."""
|
|
assert abbreviate_api_key("sk-1234") == "sk-..."
|
|
assert abbreviate_api_key("sk-test-1234") == "sk-..."
|
|
assert abbreviate_api_key("") == "sk-..."
|
|
|
|
|
|
def test_get_customer_user_header_returns_none_when_no_customer_role():
|
|
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
|
|
|
mappings = [{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}]
|
|
result = get_customer_user_header_from_mapping(mappings)
|
|
assert result is None
|
|
|
|
|
|
def test_get_customer_user_header_returns_none_for_single_non_customer_mapping():
|
|
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
|
|
|
mapping = {"header_name": "X-Only-Internal", "litellm_user_role": "internal_user"}
|
|
result = get_customer_user_header_from_mapping(mapping)
|
|
assert result is None
|
|
|
|
|
|
def test_get_customer_user_header_from_mapping_returns_customer_header():
|
|
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
|
|
|
mappings = [
|
|
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
|
|
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
|
|
]
|
|
result = get_customer_user_header_from_mapping(mappings)
|
|
assert result == ["x-openwebui-user-email"]
|
|
|
|
|
|
def test_get_customer_user_header_returns_customers_header_in_config_order_when_multiple_exist():
|
|
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
|
|
|
mappings = [
|
|
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
|
|
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
|
|
{"header_name": "X-User-Id", "litellm_user_role": "customer"},
|
|
]
|
|
result = get_customer_user_header_from_mapping(mappings)
|
|
assert result == ["x-openwebui-user-email", "x-user-id"]
|
|
|
|
|
|
def test_get_end_user_id_returns_id_from_user_header_mappings():
|
|
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
|
|
|
mappings = [
|
|
{"header_name": "x-openwebui-user-id", "litellm_user_role": "internal_user"},
|
|
{"header_name": "x-openwebui-user-email", "litellm_user_role": "customer"},
|
|
]
|
|
general_settings = {"user_header_mappings": mappings}
|
|
headers = {"x-openwebui-user-email": "1234"}
|
|
|
|
with (
|
|
patch(
|
|
"litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers",
|
|
return_value=None,
|
|
),
|
|
patch("litellm.proxy.proxy_server.general_settings", general_settings),
|
|
):
|
|
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
|
|
|
assert result == "1234"
|
|
|
|
|
|
def test_get_end_user_id_returns_first_customer_header_when_multiple_mappings_exist():
|
|
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
|
|
|
mappings = [
|
|
{"header_name": "x-openwebui-user-id", "litellm_user_role": "internal_user"},
|
|
{"header_name": "x-user-id", "litellm_user_role": "customer"},
|
|
{"header_name": "x-openwebui-user-email", "litellm_user_role": "customer"},
|
|
]
|
|
general_settings = {"user_header_mappings": mappings}
|
|
headers = {
|
|
"x-user-id": "user-456",
|
|
"x-openwebui-user-email": "user@example.com",
|
|
}
|
|
|
|
with (
|
|
patch(
|
|
"litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers",
|
|
return_value=None,
|
|
),
|
|
patch("litellm.proxy.proxy_server.general_settings", general_settings),
|
|
):
|
|
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
|
|
|
assert result == "user-456"
|
|
|
|
|
|
def test_get_end_user_id_returns_none_when_no_customer_role_in_mappings():
|
|
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
|
|
|
mappings = [
|
|
{"header_name": "x-openwebui-user-id", "litellm_user_role": "internal_user"},
|
|
]
|
|
general_settings = {"user_header_mappings": mappings}
|
|
headers = {"x-openwebui-user-id": "user-789"}
|
|
|
|
with (
|
|
patch(
|
|
"litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers",
|
|
return_value=None,
|
|
),
|
|
patch("litellm.proxy.proxy_server.general_settings", general_settings),
|
|
):
|
|
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
|
|
|
assert result is None
|
|
|
|
|
|
def test_get_end_user_id_falls_back_to_deprecated_user_header_name():
|
|
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
|
|
|
general_settings = {"user_header_name": "x-custom-user-id"}
|
|
headers = {"x-custom-user-id": "user-legacy"}
|
|
|
|
with (
|
|
patch(
|
|
"litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers",
|
|
return_value=None,
|
|
),
|
|
patch("litellm.proxy.proxy_server.general_settings", general_settings),
|
|
):
|
|
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
|
|
|
assert result == "user-legacy"
|
|
|
|
|
|
class TestCoerceUserIdToStr:
|
|
"""Unit tests for the _coerce_user_id_to_str helper."""
|
|
|
|
def test_plain_string_is_returned_verbatim(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
assert _coerce_user_id_to_str("alice@example.com") == "alice@example.com"
|
|
|
|
def test_string_is_stripped(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
assert _coerce_user_id_to_str(" bob ") == "bob"
|
|
|
|
def test_codex_opaque_identifier_is_preserved(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
codex_id = (
|
|
"user_8a4a360c36621665b341e06fb76041d9b6def732bb183eea148d4abc9d97c1de"
|
|
"_account__session_a2bce4a5-8887-44ef-b491-fbf0a55c6569"
|
|
)
|
|
assert _coerce_user_id_to_str(codex_id) == codex_id
|
|
|
|
def test_none_returns_none(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
assert _coerce_user_id_to_str(None) is None
|
|
|
|
def test_empty_string_returns_none(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
assert _coerce_user_id_to_str("") is None
|
|
assert _coerce_user_id_to_str(" ") is None
|
|
|
|
def test_dict_returns_none(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
payload = {
|
|
"device_id": "abc",
|
|
"account_uuid": "",
|
|
"session_id": "c284b8cb",
|
|
}
|
|
assert _coerce_user_id_to_str(payload) is None
|
|
|
|
def test_list_returns_none(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
assert _coerce_user_id_to_str(["a", "b"]) is None
|
|
|
|
def test_json_encoded_dict_string_passes_through_by_default(self):
|
|
"""JSON-encoded dict strings are preserved unless opt-in flag is on.
|
|
|
|
This preserves backwards compatibility: existing deployments that
|
|
intentionally pass JSON-encoded user identifiers keep working.
|
|
"""
|
|
import litellm
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
blob = (
|
|
'{"device_id":"d5abe9199ee7759a0558974e9371e78c7b38d7621aae26d6609c1de61af6afb0",'
|
|
'"account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}'
|
|
)
|
|
original = litellm.validate_end_user_id_in_db
|
|
litellm.validate_end_user_id_in_db = False
|
|
try:
|
|
assert _coerce_user_id_to_str(blob) == blob
|
|
finally:
|
|
litellm.validate_end_user_id_in_db = original
|
|
|
|
def test_json_encoded_dict_string_returns_none_when_validation_enabled(self):
|
|
import litellm
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
# Same broken shape we saw in spend logs, but pre-stringified to JSON.
|
|
blob = (
|
|
'{"device_id":"d5abe9199ee7759a0558974e9371e78c7b38d7621aae26d6609c1de61af6afb0",'
|
|
'"account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}'
|
|
)
|
|
original = litellm.validate_end_user_id_in_db
|
|
litellm.validate_end_user_id_in_db = True
|
|
try:
|
|
assert _coerce_user_id_to_str(blob) is None
|
|
finally:
|
|
litellm.validate_end_user_id_in_db = original
|
|
|
|
def test_json_encoded_list_string_passes_through_by_default(self):
|
|
import litellm
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
original = litellm.validate_end_user_id_in_db
|
|
litellm.validate_end_user_id_in_db = False
|
|
try:
|
|
assert _coerce_user_id_to_str('["a","b"]') == '["a","b"]'
|
|
finally:
|
|
litellm.validate_end_user_id_in_db = original
|
|
|
|
def test_json_encoded_list_string_returns_none_when_validation_enabled(self):
|
|
import litellm
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
original = litellm.validate_end_user_id_in_db
|
|
litellm.validate_end_user_id_in_db = True
|
|
try:
|
|
assert _coerce_user_id_to_str('["a","b"]') is None
|
|
finally:
|
|
litellm.validate_end_user_id_in_db = original
|
|
|
|
def test_int_returns_str(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
assert _coerce_user_id_to_str(12345) == "12345"
|
|
|
|
def test_bool_returns_none(self):
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
# bool is an int subclass — reject explicitly, never produce "True"/"False".
|
|
assert _coerce_user_id_to_str(True) is None
|
|
assert _coerce_user_id_to_str(False) is None
|
|
|
|
def test_brace_string_that_isnt_json_is_kept(self):
|
|
"""A string starting with `{` but failing to parse stays as-is."""
|
|
from litellm.proxy.auth.auth_utils import _coerce_user_id_to_str
|
|
|
|
assert _coerce_user_id_to_str("{not json") == "{not json"
|
|
|
|
|
|
class TestGetEndUserIdDropsMalformedBodyValues:
|
|
"""Tests that get_end_user_id_from_request_body drops dict-shaped values
|
|
rather than stringifying them into spend logs."""
|
|
|
|
def test_dict_user_falls_through_to_litellm_metadata(self):
|
|
request_body = {
|
|
"user": {
|
|
"device_id": "abc",
|
|
"session_id": "c284b8cb",
|
|
},
|
|
"litellm_metadata": {"user": "alice@example.com"},
|
|
}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result == "alice@example.com"
|
|
|
|
def test_dict_user_with_no_other_sources_returns_none(self):
|
|
request_body = {
|
|
"user": {"device_id": "abc", "session_id": "xyz"},
|
|
}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result is None
|
|
|
|
def test_json_encoded_user_string_passes_through_by_default(self):
|
|
"""JSON-encoded user strings pass through unless validation is opted in.
|
|
|
|
Gating behind ``litellm.validate_end_user_id_in_db`` keeps existing
|
|
deployments that send JSON-encoded identifiers working until they
|
|
explicitly opt into the stricter extraction.
|
|
"""
|
|
import litellm
|
|
|
|
blob = '{"device_id":"d5abe9199ee7759a","account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}'
|
|
request_body = {"user": blob}
|
|
|
|
original = litellm.validate_end_user_id_in_db
|
|
litellm.validate_end_user_id_in_db = False
|
|
try:
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
finally:
|
|
litellm.validate_end_user_id_in_db = original
|
|
|
|
assert result == blob
|
|
|
|
def test_json_encoded_user_string_returns_none_when_validation_enabled(self):
|
|
import litellm
|
|
|
|
request_body = {
|
|
"user": (
|
|
'{"device_id":"d5abe9199ee7759a","account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}'
|
|
),
|
|
}
|
|
|
|
original = litellm.validate_end_user_id_in_db
|
|
litellm.validate_end_user_id_in_db = True
|
|
try:
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
finally:
|
|
litellm.validate_end_user_id_in_db = original
|
|
|
|
assert result is None
|
|
|
|
def test_plain_string_user_is_preserved(self):
|
|
request_body = {"user": "alice@example.com"}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result == "alice@example.com"
|
|
|
|
def test_codex_opaque_user_is_preserved(self):
|
|
codex_id = (
|
|
"user_8a4a360c36621665b341e06fb76041d9b6def732bb183eea148d4abc9d97c1de"
|
|
"_account__session_a2bce4a5-8887-44ef-b491-fbf0a55c6569"
|
|
)
|
|
request_body = {"user": codex_id}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result == codex_id
|
|
|
|
def test_int_user_is_coerced_to_string(self):
|
|
request_body = {"user": 12345}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result == "12345"
|
|
|
|
def test_list_user_falls_through(self):
|
|
request_body = {
|
|
"user": ["a", "b"],
|
|
"safety_identifier": "alice@example.com",
|
|
}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result == "alice@example.com"
|
|
|
|
def test_dict_safety_identifier_returns_none(self):
|
|
request_body = {
|
|
"safety_identifier": {"device_id": "abc"},
|
|
}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result is None
|
|
|
|
def test_dict_metadata_user_id_returns_none(self):
|
|
request_body = {
|
|
"metadata": {"user_id": {"device_id": "abc"}},
|
|
}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result is None
|
|
|
|
def test_whitespace_user_falls_through(self):
|
|
request_body = {"user": " ", "safety_identifier": "alice@example.com"}
|
|
|
|
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers={})
|
|
|
|
assert result == "alice@example.com"
|
|
|
|
def test_dict_user_header_falls_through_to_body(self):
|
|
"""A dict-shaped value in a configured user-id header is dropped, not stringified."""
|
|
general_settings = {"user_header_name": "x-custom-user-id"}
|
|
# A header value will normally be a str, but be defensive: the coercion
|
|
# must drop anything that isn't a usable identifier.
|
|
headers = {"x-custom-user-id": {"device_id": "abc"}}
|
|
request_body = {"user": "alice@example.com"}
|
|
|
|
with (
|
|
patch(
|
|
"litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers",
|
|
return_value=None,
|
|
),
|
|
patch("litellm.proxy.proxy_server.general_settings", general_settings),
|
|
):
|
|
result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers)
|
|
|
|
assert result == "alice@example.com"
|
|
|
|
|
|
def _make_deployment_dict(model_name: str, tpm: Optional[int] = None, rpm: Optional[int] = None) -> dict:
|
|
"""Helper to build a minimal deployment dict as returned by router.get_model_list."""
|
|
litellm_params: dict = {"model": model_name}
|
|
if tpm is not None:
|
|
litellm_params["default_api_key_tpm_limit"] = tpm
|
|
if rpm is not None:
|
|
litellm_params["default_api_key_rpm_limit"] = rpm
|
|
return {"model_name": model_name, "litellm_params": litellm_params}
|
|
|
|
|
|
_ROUTER_PATCH = "litellm.proxy.proxy_server.llm_router"
|
|
|
|
|
|
class TestDeploymentDefaultRpmLimit:
|
|
"""Tests for deployment default_api_key_rpm_limit fallback in get_key_model_rpm_limit."""
|
|
|
|
def test_returns_deployment_default_when_key_has_no_limits(self):
|
|
"""Case 2 from spec: key has no model-specific limits, falls back to deployment default."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 200}
|
|
|
|
def test_key_model_limit_takes_priority_over_deployment_default(self):
|
|
"""Case 1 from spec: key model-specific limit wins over deployment default."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"model_rpm_limit": {"model1": 10}},
|
|
)
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 10}
|
|
|
|
def test_returns_none_when_no_deployment_default_and_no_key_limits(self):
|
|
"""Returns None when neither the key nor the deployment has any rpm limit."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
_make_deployment_dict("model1") # no rpm default
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result is None
|
|
|
|
def test_returns_none_without_model_name_even_when_deployment_has_default(self):
|
|
"""No model_name means deployment fallback is skipped."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_rpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
def test_returns_none_when_llm_router_is_none(self):
|
|
"""No router means deployment fallback returns None gracefully."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
with patch(_ROUTER_PATCH, None):
|
|
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result is None
|
|
|
|
def test_returns_minimum_across_multiple_deployments(self):
|
|
"""When multiple deployments share a model name, the minimum rpm limit is used."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
_make_deployment_dict("model1", rpm=200),
|
|
_make_deployment_dict("model1", rpm=50),
|
|
_make_deployment_dict("model1", rpm=150),
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 50}
|
|
|
|
def test_ignores_deployments_without_default_when_others_have_it(self):
|
|
"""Deployments missing the field are skipped; min is taken over those that have it."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
_make_deployment_dict("model1"), # no rpm default
|
|
_make_deployment_dict("model1", rpm=75),
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 75}
|
|
|
|
def test_skips_deployments_with_malformed_limit_value(self):
|
|
"""Deployments with non-integer-parseable limit values are skipped without raising."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
{
|
|
"model_name": "model1",
|
|
"litellm_params": {"default_api_key_rpm_limit": "not-a-number"},
|
|
},
|
|
_make_deployment_dict("model1", rpm=100),
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
|
|
# The malformed deployment is skipped; the valid one provides 100
|
|
assert result == {"model1": 100}
|
|
|
|
|
|
class TestDeploymentDefaultTpmLimit:
|
|
"""Tests for deployment default_api_key_tpm_limit fallback in get_key_model_tpm_limit."""
|
|
|
|
def test_returns_deployment_default_when_key_has_no_limits(self):
|
|
"""Case 2 from spec: key has no model-specific limits, falls back to deployment default."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 100}
|
|
|
|
def test_key_model_limit_takes_priority_over_deployment_default(self):
|
|
"""Case 1 from spec: key model-specific limit wins over deployment default."""
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
metadata={"model_tpm_limit": {"model1": 20}},
|
|
)
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 20}
|
|
|
|
def test_returns_none_when_no_deployment_default_and_no_key_limits(self):
|
|
"""Returns None when neither the key nor the deployment has any tpm limit."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
_make_deployment_dict("model1") # no tpm default
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result is None
|
|
|
|
def test_returns_none_without_model_name_even_when_deployment_has_default(self):
|
|
"""No model_name means deployment fallback is skipped."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_tpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
def test_returns_none_when_llm_router_is_none(self):
|
|
"""No router means deployment fallback returns None gracefully."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
with patch(_ROUTER_PATCH, None):
|
|
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result is None
|
|
|
|
def test_returns_minimum_across_multiple_deployments(self):
|
|
"""When multiple deployments share a model name, the minimum tpm limit is used."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
_make_deployment_dict("model1", tpm=1000),
|
|
_make_deployment_dict("model1", tpm=300),
|
|
_make_deployment_dict("model1", tpm=700),
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 300}
|
|
|
|
def test_ignores_deployments_without_default_when_others_have_it(self):
|
|
"""Deployments missing the field are skipped; min is taken over those that have it."""
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
mock_router = MagicMock()
|
|
mock_router.get_model_list.return_value = [
|
|
_make_deployment_dict("model1"), # no tpm default
|
|
_make_deployment_dict("model1", tpm=400),
|
|
]
|
|
with patch(_ROUTER_PATCH, mock_router):
|
|
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
|
|
assert result == {"model1": 400}
|
|
|
|
|
|
class TestGetProjectModelRpmLimit:
|
|
"""Tests for get_project_model_rpm_limit function."""
|
|
|
|
def test_returns_project_metadata_rpm_limit(self):
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
project_metadata={"model_rpm_limit": {"gpt-4": 200}},
|
|
)
|
|
result = get_project_model_rpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 200}
|
|
|
|
def test_returns_none_when_no_project_metadata(self):
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
result = get_project_model_rpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
def test_returns_none_when_project_metadata_missing_key(self):
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
project_metadata={"other_key": "value"},
|
|
)
|
|
result = get_project_model_rpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
|
|
class TestGetProjectModelTpmLimit:
|
|
"""Tests for get_project_model_tpm_limit function."""
|
|
|
|
def test_returns_project_metadata_tpm_limit(self):
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
project_metadata={"model_tpm_limit": {"gpt-4": 50000}},
|
|
)
|
|
result = get_project_model_tpm_limit(user_api_key_dict)
|
|
assert result == {"gpt-4": 50000}
|
|
|
|
def test_returns_none_when_no_project_metadata(self):
|
|
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
|
|
result = get_project_model_tpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
def test_returns_none_when_project_metadata_missing_key(self):
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
api_key="sk-123",
|
|
project_metadata={"other_key": "value"},
|
|
)
|
|
result = get_project_model_tpm_limit(user_api_key_dict)
|
|
assert result is None
|
|
|
|
|
|
class TestCheckCompleteCredentials:
|
|
"""Tests for the api_key validation in check_complete_credentials."""
|
|
|
|
def test_returns_false_when_api_key_missing(self):
|
|
result = check_complete_credentials({"model": "gpt-4"})
|
|
assert result is False
|
|
|
|
def test_returns_false_when_api_key_is_none(self):
|
|
result = check_complete_credentials({"model": "gpt-4", "api_key": None})
|
|
assert result is False
|
|
|
|
def test_returns_false_when_api_key_is_empty_string(self):
|
|
result = check_complete_credentials({"model": "gpt-4", "api_key": ""})
|
|
assert result is False
|
|
|
|
def test_returns_false_when_api_key_is_whitespace(self):
|
|
result = check_complete_credentials({"model": "gpt-4", "api_key": " "})
|
|
assert result is False
|
|
|
|
def test_returns_true_when_api_key_is_valid(self):
|
|
result = check_complete_credentials({"model": "gpt-4", "api_key": "sk-valid"})
|
|
assert result is True
|
|
|
|
|
|
class TestCheckCompleteCredentialsBlocksSSRF:
|
|
"""
|
|
Even with credentials supplied, ``api_base`` / ``base_url`` must not
|
|
point at private / internal / cloud-metadata addresses. Without this
|
|
the gate accepts ``api_key=anything`` plus a malicious target and the
|
|
proxy is used as an SSRF pivot.
|
|
|
|
The check only runs when ``litellm.user_url_validation`` is True, so
|
|
every test in this class flips the toggle. Tests stay mock-only — no
|
|
real DNS is performed.
|
|
"""
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _enable_url_validation(self, monkeypatch):
|
|
import litellm
|
|
|
|
monkeypatch.setattr(litellm, "user_url_validation", True, raising=False)
|
|
|
|
@pytest.mark.parametrize(
|
|
"url_field",
|
|
["api_base", "base_url"],
|
|
)
|
|
@pytest.mark.parametrize(
|
|
"blocked_url",
|
|
[
|
|
"http://169.254.169.254/latest/meta-data/iam/security-credentials/",
|
|
"http://metadata.google.internal/computeMetadata/v1/",
|
|
"http://127.0.0.1:8080/admin",
|
|
"http://10.0.0.1/",
|
|
"http://192.168.1.1/",
|
|
],
|
|
)
|
|
def test_rejects_private_or_metadata_targets(self, url_field, blocked_url):
|
|
from litellm.litellm_core_utils.url_utils import SSRFError
|
|
|
|
with patch(
|
|
"litellm.proxy.auth.auth_utils.validate_url",
|
|
side_effect=SSRFError(f"blocked: {blocked_url}"),
|
|
):
|
|
with pytest.raises(ValueError, match="is rejected by the SSRF guard") as exc_info:
|
|
check_complete_credentials(
|
|
{
|
|
"model": "gpt-4",
|
|
"api_key": "sk-some-clientside-key",
|
|
url_field: blocked_url,
|
|
}
|
|
)
|
|
assert url_field in str(exc_info.value)
|
|
assert "SSRF" in str(exc_info.value)
|
|
|
|
def test_allows_public_target_when_validate_url_passes(self):
|
|
# ``validate_url`` is mocked so no real DNS is performed.
|
|
with patch(
|
|
"litellm.proxy.auth.auth_utils.validate_url",
|
|
return_value=("https://api.openai.com/v1", "api.openai.com"),
|
|
):
|
|
result = check_complete_credentials(
|
|
{
|
|
"model": "gpt-4",
|
|
"api_key": "sk-some-clientside-key",
|
|
"api_base": "https://api.openai.com/v1",
|
|
}
|
|
)
|
|
assert result is True
|
|
|
|
def test_skips_url_validation_when_toggle_is_off(self, monkeypatch):
|
|
# Admins who disable ``user_url_validation`` (default) should not
|
|
# have requests rejected at the proxy boundary even if the URL
|
|
# would fail the SSRF guard.
|
|
import litellm
|
|
|
|
monkeypatch.setattr(litellm, "user_url_validation", False, raising=False)
|
|
with patch(
|
|
"litellm.proxy.auth.auth_utils.validate_url",
|
|
) as mocked:
|
|
result = check_complete_credentials(
|
|
{
|
|
"model": "gpt-4",
|
|
"api_key": "sk-some-clientside-key",
|
|
"api_base": "http://127.0.0.1:8080/admin",
|
|
}
|
|
)
|
|
assert result is True
|
|
mocked.assert_not_called()
|
|
|
|
|
|
class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride:
|
|
"""
|
|
When the caller redirects ``api_base`` / ``base_url`` to their own
|
|
server, admin-set fields like ``OpenAI-Organization``, ``extra_body``,
|
|
AWS / Vertex / Azure tokens, and per-deployment ``api_version`` must
|
|
NOT flow through to that destination.
|
|
"""
|
|
|
|
def test_clears_admin_organization_and_extra_body_on_base_override(self):
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
admin_params = {
|
|
"model": "gpt-4",
|
|
"api_key": "sk-admin-key",
|
|
"api_base": "https://admin.upstream/v1",
|
|
"organization": "org-admin-corp",
|
|
"extra_body": {"x-admin-secret": "super-secret"},
|
|
"api_version": "2026-04-01",
|
|
}
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params=dict(admin_params),
|
|
request_kwargs={
|
|
"api_key": "sk-attacker",
|
|
"api_base": "https://attacker.example",
|
|
},
|
|
)
|
|
assert out["api_base"] == "https://attacker.example"
|
|
assert out["api_key"] == "sk-attacker"
|
|
assert "organization" not in out
|
|
assert "extra_body" not in out
|
|
assert "api_version" not in out
|
|
|
|
def test_clears_aws_and_vertex_secrets_on_base_override(self):
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
admin_params = {
|
|
"model": "bedrock/claude-3",
|
|
"aws_access_key_id": "AKIA-EXAMPLE",
|
|
"aws_secret_access_key": "secret-example",
|
|
"aws_session_token": "session-example",
|
|
"vertex_credentials": '{"private_key":"-----BEGIN..."}',
|
|
"vertex_project": "admin-gcp-project",
|
|
}
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params=dict(admin_params),
|
|
request_kwargs={"base_url": "https://attacker.example", "api_key": "sk-caller"},
|
|
)
|
|
assert "aws_access_key_id" not in out
|
|
assert "aws_secret_access_key" not in out
|
|
assert "aws_session_token" not in out
|
|
assert "vertex_credentials" not in out
|
|
assert "vertex_project" not in out
|
|
|
|
def test_clears_nvcf_function_id_on_base_override(self):
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
admin_params = {
|
|
"model": "nvidia_riva/parakeet",
|
|
"api_base": "grpc.nvcf.nvidia.com:443",
|
|
"api_key": "nvapi-admin",
|
|
"nvcf_function_id": "admin-pinned-function",
|
|
}
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params=dict(admin_params),
|
|
request_kwargs={"api_base": "self-hosted.example.com:50051", "api_key": "sk-caller"},
|
|
)
|
|
assert out["api_base"] == "self-hosted.example.com:50051"
|
|
assert "nvcf_function_id" not in out
|
|
|
|
def test_clears_use_ssl_on_base_override(self):
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
admin_params = {
|
|
"model": "nvidia_riva/parakeet",
|
|
"api_base": "grpc.nvcf.nvidia.com:443",
|
|
"api_key": "nvapi-admin",
|
|
"use_ssl": True,
|
|
}
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params=dict(admin_params),
|
|
request_kwargs={"api_base": "self-hosted.example.com:50051", "api_key": "sk-caller"},
|
|
)
|
|
assert out["api_base"] == "self-hosted.example.com:50051"
|
|
assert "use_ssl" not in out
|
|
|
|
def test_caller_resupplied_value_overrides_admin_value_on_base_override(self):
|
|
# When the caller redirects ``api_base`` and *also* supplies their
|
|
# own value for one of the admin fields (e.g. ``organization``),
|
|
# the caller's value must win — never the admin's. The naive
|
|
# ``if field not in request_kwargs: pop`` shape lets a caller echo
|
|
# the field name with any value (or empty string) to keep the
|
|
# admin's value forwarded, which is the exfiltration vector this
|
|
# test guards against.
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params={
|
|
"api_base": "https://admin.upstream/v1",
|
|
"organization": "org-admin",
|
|
"extra_body": {"admin": "value"},
|
|
},
|
|
request_kwargs={
|
|
"api_base": "https://attacker.example",
|
|
"api_key": "sk-caller",
|
|
"organization": "org-attacker",
|
|
"extra_body": {"attacker": "value"},
|
|
},
|
|
)
|
|
assert out["organization"] == "org-attacker"
|
|
assert out["extra_body"] == {"attacker": "value"}
|
|
|
|
def test_field_echo_does_not_preserve_admin_value(self):
|
|
# Regression: a caller that echoes an admin-config field name with
|
|
# an *empty* value (or any value) must not be able to keep the
|
|
# admin's value in ``litellm_params``.
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params={
|
|
"api_base": "https://admin.upstream/v1",
|
|
"organization": "org-admin-secret",
|
|
"extra_body": {"x-admin-only": "secret"},
|
|
},
|
|
request_kwargs={
|
|
"api_base": "https://attacker.example",
|
|
"api_key": "sk-caller",
|
|
"organization": "",
|
|
"extra_body": "",
|
|
},
|
|
)
|
|
assert out["organization"] == ""
|
|
assert out["extra_body"] == ""
|
|
assert "org-admin-secret" not in str(out)
|
|
|
|
def test_no_clearing_when_only_api_key_overridden(self):
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
# Caller only overrides api_key (BYOK pattern); admin's organization /
|
|
# extra_body / region still apply because the destination is unchanged.
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params={
|
|
"api_base": "https://admin.upstream/v1",
|
|
"organization": "org-admin",
|
|
"api_version": "2026-04-01",
|
|
},
|
|
request_kwargs={"api_key": "sk-byok"},
|
|
)
|
|
assert out["organization"] == "org-admin"
|
|
assert out["api_version"] == "2026-04-01"
|
|
assert out["api_base"] == "https://admin.upstream/v1"
|
|
|
|
def test_client_api_key_used_when_supplied_with_base_override(self):
|
|
from litellm.router_utils.clientside_credential_handler import (
|
|
get_dynamic_litellm_params,
|
|
)
|
|
|
|
out = get_dynamic_litellm_params(
|
|
litellm_params={
|
|
"model": "gpt-4",
|
|
"api_key": "sk-admin-secret",
|
|
"api_base": "https://admin.upstream/v1",
|
|
},
|
|
request_kwargs={
|
|
"api_base": "https://attacker.example",
|
|
"api_key": "sk-client-byok",
|
|
},
|
|
)
|
|
assert out["api_key"] == "sk-client-byok"
|
|
assert "sk-admin-secret" not in str(out)
|
|
|
|
|
|
_OPENAI_CHAT_RESPONSE = {
|
|
"id": "chatcmpl-x",
|
|
"object": "chat.completion",
|
|
"created": 1,
|
|
"model": "gpt-4",
|
|
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
|
|
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
|
}
|
|
|
|
|
|
class TestClientsideBaseOverrideOutboundKey:
|
|
"""Drive a completion through the router and assert on the outbound request
|
|
when the caller overrides ``api_base``."""
|
|
|
|
def _router(self):
|
|
from litellm import Router
|
|
|
|
return Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-4",
|
|
"litellm_params": {
|
|
"model": "openai/gpt-4",
|
|
"api_key": "sk-SERVER-CONFIG",
|
|
"api_base": "https://admin.upstream/v1",
|
|
},
|
|
}
|
|
]
|
|
)
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _ambient_server_key(self, monkeypatch):
|
|
import litellm
|
|
|
|
monkeypatch.setenv("OPENAI_API_KEY", "sk-SERVER-ENV")
|
|
monkeypatch.setattr(litellm, "api_key", None, raising=False)
|
|
|
|
def test_caller_key_override_sends_caller_key_never_server_key(self):
|
|
import httpx
|
|
import respx
|
|
|
|
with respx.mock:
|
|
route = respx.post("https://caller.example/v1/chat/completions").mock(
|
|
return_value=httpx.Response(200, json=_OPENAI_CHAT_RESPONSE)
|
|
)
|
|
self._router().completion(
|
|
model="gpt-4",
|
|
messages=[{"role": "user", "content": "hi"}],
|
|
api_base="https://caller.example/v1",
|
|
api_key="sk-CALLER",
|
|
)
|
|
authorization = route.calls.last.request.headers.get("authorization")
|
|
assert authorization == "Bearer sk-CALLER"
|
|
assert "SERVER" not in (authorization or "")
|
|
|
|
|
|
def _rounds_deep_api_base_payload(rounds, field):
|
|
"""Build a fallbacks payload with ``api_base`` on a target nested ``rounds``
|
|
fallback-rounds deep, each round wrapped in its own grouping dict."""
|
|
node = {"model": "leaf", "api_base": "https://attacker.example"}
|
|
for i in range(rounds):
|
|
node = {"model": f"m{i}", field: [{"grp": [node]}]}
|
|
return {"model": "gpt-4", field: [{"grp": [node]}]}
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksFallbackSmuggle:
|
|
"""``is_request_body_safe`` runs the banned-param check on every dict target
|
|
inside the fallback lists."""
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _disable_url_validation(self, monkeypatch):
|
|
import litellm
|
|
|
|
monkeypatch.setattr(litellm, "user_url_validation", False, raising=False)
|
|
|
|
@pytest.mark.parametrize(
|
|
"fallback_key",
|
|
["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"],
|
|
)
|
|
def test_api_base_smuggled_via_nested_fallback_is_rejected(self, fallback_key):
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
fallback_key: [
|
|
{
|
|
"gpt-4": [
|
|
{"model": "evil", "api_base": "https://attacker.example"},
|
|
]
|
|
}
|
|
],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_string_only_fallbacks_are_accepted(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"fallbacks": [{"gpt-4": ["gpt-3.5-turbo", "claude-3-haiku"]}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_benign_dict_fallback_entry_is_accepted(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"fallbacks": [{"gpt-4": [{"model": "gpt-3.5-turbo"}]}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_smuggled_fallback_allowed_under_proxy_wide_opt_in(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"fallbacks": [{"gpt-4": [{"model": "byok", "api_base": "https://my-byok.example"}]}],
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
@pytest.mark.parametrize(
|
|
"fallback_field",
|
|
["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"],
|
|
)
|
|
@pytest.mark.parametrize("surface", ["top_level", "router_settings_override"])
|
|
def test_deeply_nested_api_base_smuggle_rejected_on_both_surfaces(self, fallback_field, surface):
|
|
nested = [
|
|
{
|
|
"always-fail": [
|
|
{
|
|
"model": "x",
|
|
fallback_field: [{"x": [{"model": "deepseek-chat", "api_base": "http://attacker"}]}],
|
|
}
|
|
]
|
|
}
|
|
]
|
|
request_body = {"model": "gpt-4"}
|
|
if surface == "top_level":
|
|
request_body[fallback_field] = nested
|
|
else:
|
|
request_body["router_settings_override"] = {fallback_field: nested}
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body=request_body,
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_router_settings_override_single_level_api_base_rejected(self):
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"router_settings_override": {
|
|
"fallbacks": [{"gpt-4": [{"model": "x", "api_base": "http://attacker"}]}]
|
|
},
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_model_less_config_dict_api_base_rejected(self):
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"fallbacks": [{"gpt-4": [{"api_base": "http://attacker"}]}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_nested_api_base_caught_across_router_fallback_rounds(self):
|
|
"""An ``api_base`` target nested ``ROUTER_MAX_FALLBACKS - 1`` rounds deep
|
|
is still reached and rejected."""
|
|
import litellm
|
|
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body=_rounds_deep_api_base_payload(litellm.ROUTER_MAX_FALLBACKS - 1, "fallbacks"),
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_grouping_only_deep_chain_is_rejected_at_depth_limit(self):
|
|
"""A deep grouping-only chain (``{"g": [{"g": [...]}]}``) is rejected at the
|
|
validation-depth limit rather than accepted or raising RecursionError."""
|
|
node: object = ["safe-model"]
|
|
for _ in range(5000):
|
|
node = [{"grp": node}]
|
|
with pytest.raises(ValueError, match="depth"):
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "fallbacks": node},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_pathologically_deep_model_nesting_is_rejected(self):
|
|
with pytest.raises(ValueError, match="depth"):
|
|
is_request_body_safe(
|
|
request_body=_rounds_deep_api_base_payload(5000, "fallbacks"),
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeRejectsUrlValuedFallback:
|
|
@pytest.mark.parametrize("fallback_field", ["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"])
|
|
def test_url_valued_string_fallback_is_rejected(self, fallback_field):
|
|
with pytest.raises(ValueError, match="URL-valued fallback"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
fallback_field: [{"gpt-4": ["huggingface/http://attacker.example/path"]}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
@pytest.mark.parametrize("fallback_field", ["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"])
|
|
def test_url_valued_dict_model_fallback_is_rejected(self, fallback_field):
|
|
with pytest.raises(ValueError, match="URL-valued fallback"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
fallback_field: [{"gpt-4": [{"model": "huggingface/http://attacker.example/path"}]}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_ordinary_string_fallback_is_allowed(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "fallbacks": [{"gpt-4": ["gpt-4-backup"]}]},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_ordinary_dict_model_fallback_is_allowed(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "fallbacks": [{"gpt-4": [{"model": "gpt-4-backup"}]}]},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksEndpointTargetingFields:
|
|
"""
|
|
``is_request_body_safe`` rejects request-body fields that retarget the
|
|
outbound request to a caller-controlled host. Beyond the original
|
|
``api_base`` / ``base_url``, the same protection must apply to:
|
|
|
|
* ``aws_bedrock_runtime_endpoint`` — Bedrock endpoint redirect; an
|
|
attacker-controlled value coerces the proxy to authenticate against
|
|
their host with the admin's AWS creds.
|
|
* ``langsmith_base_url`` — Langsmith callback host; attacker-controlled
|
|
values exfiltrate the entire request payload (incl. message content)
|
|
via the observability hook.
|
|
* ``langfuse_host`` — same exfil vector via the Langfuse hook.
|
|
"""
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _disable_url_validation(self, monkeypatch):
|
|
# The new banned-params entries should be rejected even when
|
|
# ``user_url_validation`` is off — the gate isn't the URL guard,
|
|
# it's the banned-params list.
|
|
import litellm
|
|
|
|
monkeypatch.setattr(litellm, "user_url_validation", False, raising=False)
|
|
|
|
@pytest.mark.parametrize(
|
|
"field",
|
|
[
|
|
"aws_bedrock_runtime_endpoint",
|
|
"langsmith_base_url",
|
|
"langfuse_host",
|
|
"posthog_host",
|
|
"braintrust_host",
|
|
"slack_webhook_url",
|
|
"s3_endpoint_url",
|
|
"sagemaker_base_url",
|
|
"deployment_url",
|
|
],
|
|
)
|
|
def test_endpoint_targeting_field_in_request_body_is_rejected(self, field):
|
|
with pytest.raises(ValueError, match="Rejected Request") as exc:
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", field: "https://attacker.example"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
# The function lists the offending param name in the error.
|
|
assert field in str(exc.value)
|
|
|
|
@pytest.mark.parametrize(
|
|
"field",
|
|
["api_base", "base_url", "user_config", "langfuse_host", "slack_webhook_url"],
|
|
)
|
|
def test_api_key_does_not_bypass_blocklist(self, field):
|
|
# Regression: the historical ``check_complete_credentials`` clause
|
|
# made the entire blocklist a no-op for any caller that supplied
|
|
# a non-empty ``api_key``. That bypass turned every missing entry
|
|
# on the blocklist into an SSRF / credential-exfil hole. Verify
|
|
# that supplying an api_key (alongside the banned param) does NOT
|
|
# bypass the gate — it can only be opened by an admin opt-in.
|
|
with pytest.raises(ValueError, match="Rejected Request") as exc:
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"api_key": "sk-anything",
|
|
field: "https://attacker.example",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
assert field in str(exc.value)
|
|
|
|
def test_admin_opt_in_proxy_wide_still_allows(self):
|
|
# ``general_settings.allow_client_side_credentials = True`` remains
|
|
# the documented proxy-wide BYOK opt-in.
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "api_base": "https://my-byok.example"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksBedrockProjectOverride:
|
|
"""``aws_bedrock_project_id`` pins a deployment to a Bedrock project so
|
|
that project's data-retention policy applies to its requests. A
|
|
caller-supplied value would run the request under any project reachable
|
|
with the deployment's shared AWS credentials, bypassing the configured
|
|
retention/accounting association."""
|
|
|
|
def test_project_id_in_request_body_is_rejected(self):
|
|
with pytest.raises(ValueError, match="aws_bedrock_project_id"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"aws_bedrock_project_id": "proj_attacker000000",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_admin_opt_in_proxy_wide_allows_project_id(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"aws_bedrock_project_id": "proj_byok000000",
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksClaudePlatformWorkspaceOverride:
|
|
@pytest.mark.parametrize(
|
|
"alias", ["workspace_id", "aws_workspace_id", "anthropic_workspace_id", "anthropic-workspace-id"]
|
|
)
|
|
def test_workspace_alias_in_request_body_is_rejected(self, alias):
|
|
with pytest.raises(ValueError, match=alias):
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", alias: "wrkspc_attacker"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_admin_opt_in_proxy_wide_allows_workspace_id(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "workspace_id": "wrkspc_byok"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
class TestIsRequestBodySafeBlocksRustOptIn:
|
|
"""``rust`` hands the whole call to the Rust core, which signs and sends
|
|
with its own HTTP client rather than the one the deployment configured, and
|
|
reports no ``post_call``. The proxy splats the request body straight into
|
|
the router, and ``rust`` is a litellm param, so it lands in
|
|
``litellm_params`` and the gate honours it: without this entry any
|
|
authenticated caller picks a transport and a callback surface the admin
|
|
never chose. It stays a deployment decision, liftable only by the same
|
|
admin opt-in as the rest of the list."""
|
|
|
|
def test_rust_in_request_body_is_rejected(self):
|
|
with pytest.raises(ValueError, match="rust"):
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "rust": True},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_rust_under_extra_body_is_rejected(self):
|
|
with pytest.raises(ValueError, match="not allowed in request body"):
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "extra_body": {"rust": True}},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_api_key_does_not_bypass_the_rust_block(self):
|
|
with pytest.raises(ValueError, match="rust"):
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "api_key": "sk-anything", "rust": True},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_admin_opt_in_proxy_wide_allows_rust(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "rust": True},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_body_without_rust_is_still_allowed(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "temperature": 0.7},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksVertexCredentialAlias:
|
|
@pytest.mark.parametrize("field", ["vertex_ai_credentials"])
|
|
def test_field_in_request_body_is_rejected(self, field):
|
|
with pytest.raises(ValueError, match=field):
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", field: "attacker-supplied"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
@pytest.mark.parametrize("field", ["vertex_ai_credentials"])
|
|
def test_admin_opt_in_proxy_wide_allows(self, field):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", field: "byok-supplied"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_legitimate_request_body_param_still_allowed(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"temperature": 0.7,
|
|
"max_tokens": 128,
|
|
"user": "end-user-123",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksNVCFFunctionOverride:
|
|
"""``nvcf_function_id`` is rejected as a request-body param unless the
|
|
admin opted in proxy-wide or per-deployment."""
|
|
|
|
def test_nvcf_function_id_in_request_body_is_rejected(self):
|
|
with pytest.raises(ValueError, match="nvcf_function_id"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"nvcf_function_id": "caller-supplied",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
|
|
def test_nvcf_function_id_with_api_key_still_rejected(self):
|
|
with pytest.raises(ValueError, match="nvcf_function_id"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"api_key": "sk-anything",
|
|
"nvcf_function_id": "caller-supplied",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
|
|
def test_admin_opt_in_proxy_wide_allows_nvcf_function_id(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"nvcf_function_id": "byok-function-id",
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_admin_opt_in_per_deployment_allows_nvcf_function_id(self, monkeypatch):
|
|
"""The error message lists per-deployment ``configurable_clientside_auth_params``
|
|
as a second opt-in. Cover that path too so it can't silently regress."""
|
|
from litellm.proxy.auth import auth_utils
|
|
|
|
monkeypatch.setattr(
|
|
auth_utils,
|
|
"_allow_model_level_clientside_configurable_parameters",
|
|
lambda model, param, request_body_value, llm_router: param == "nvcf_function_id",
|
|
)
|
|
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"nvcf_function_id": "byok-function-id",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksRivaUseSsl:
|
|
"""``use_ssl`` is rejected as a request-body param unless the admin
|
|
opted in proxy-wide or per-deployment."""
|
|
|
|
def test_use_ssl_in_request_body_is_rejected(self):
|
|
with pytest.raises(ValueError, match="use_ssl"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"use_ssl": False,
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
|
|
def test_admin_opt_in_proxy_wide_allows_use_ssl(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"use_ssl": True,
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_admin_opt_in_per_deployment_allows_use_ssl(self, monkeypatch):
|
|
from litellm.proxy.auth import auth_utils
|
|
|
|
monkeypatch.setattr(
|
|
auth_utils,
|
|
"_allow_model_level_clientside_configurable_parameters",
|
|
lambda model, param, request_body_value, llm_router: param == "use_ssl",
|
|
)
|
|
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"use_ssl": True,
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksBedrockTags:
|
|
"""``bedrock_tags`` lands as AWS resource tags on Bedrock batch jobs
|
|
created with the proxy's AWS identity, so a caller-supplied value can
|
|
forge ownership or cost-allocation labels; like
|
|
``aws_bedrock_project_id`` it is blocked without an admin opt-in."""
|
|
|
|
def test_bedrock_tags_in_request_body_is_rejected(self):
|
|
with pytest.raises(ValueError, match="bedrock_tags"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "bedrock-batch-opus",
|
|
"bedrock_tags": [{"key": "application", "value": "genai-proxy"}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="bedrock-batch-opus",
|
|
)
|
|
|
|
def test_admin_opt_in_proxy_wide_allows_bedrock_tags(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "bedrock-batch-opus",
|
|
"bedrock_tags": [{"key": "application", "value": "genai-proxy"}],
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="bedrock-batch-opus",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_admin_opt_in_per_deployment_allows_bedrock_tags(self):
|
|
from litellm import Router
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "bedrock-batch-opus",
|
|
"litellm_params": {
|
|
"model": "bedrock/us.anthropic.claude-opus-4-7",
|
|
"configurable_clientside_auth_params": ["bedrock_tags"],
|
|
},
|
|
}
|
|
]
|
|
)
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "bedrock-batch-opus",
|
|
"bedrock_tags": [{"key": "application", "value": "genai-proxy"}],
|
|
},
|
|
general_settings={},
|
|
llm_router=router,
|
|
model="bedrock-batch-opus",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_per_deployment_opt_in_for_other_param_still_rejects_bedrock_tags(self):
|
|
from litellm import Router
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "bedrock-batch-opus",
|
|
"litellm_params": {
|
|
"model": "bedrock/us.anthropic.claude-opus-4-7",
|
|
"configurable_clientside_auth_params": ["api_base"],
|
|
},
|
|
}
|
|
]
|
|
)
|
|
with pytest.raises(ValueError, match="bedrock_tags"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "bedrock-batch-opus",
|
|
"bedrock_tags": [{"key": "application", "value": "genai-proxy"}],
|
|
},
|
|
general_settings={},
|
|
llm_router=router,
|
|
model="bedrock-batch-opus",
|
|
)
|
|
|
|
|
|
# ── is_request_body_safe nested-config recursion (VERIA-6) ────────────────────
|
|
|
|
|
|
class TestIsRequestBodySafeNestedConfig:
|
|
"""The Milvus vector store transformer unpacks
|
|
``litellm_embedding_config`` as ``**kwargs`` into ``litellm.embedding(...)``
|
|
— same SSRF / credential-exfil surface as a top-level ``api_base`` in
|
|
the request body. ``is_request_body_safe`` must recurse into this
|
|
nested dict so a banned param can't be smuggled in via nesting."""
|
|
|
|
def test_root_level_api_base_blocked_when_no_opt_in(self):
|
|
"""Sanity check: pre-existing root-level enforcement still works."""
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body={"api_base": "https://attacker.example.com"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_nested_api_base_in_embedding_config_blocked(self):
|
|
"""Smuggling ``api_base`` inside ``litellm_embedding_config`` is
|
|
the VERIA-6 bypass — must be blocked by the recursive check."""
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"litellm_embedding_config": {
|
|
"api_base": "https://attacker.example.com",
|
|
"api_key": "leaked-key",
|
|
}
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="milvus-store",
|
|
)
|
|
|
|
def test_nested_nvcf_function_id_in_metadata_blocked(self):
|
|
"""Smuggling ``nvcf_function_id`` via ``metadata`` / ``extra_body``
|
|
is the same shape as the VERIA-6 ``api_base`` bypass — must be
|
|
rejected by the recursive walk so the NVCF override gate cannot
|
|
be sidestepped with nesting."""
|
|
with pytest.raises(ValueError, match="nvcf_function_id"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "nvidia_riva/parakeet",
|
|
"litellm_metadata": {"nvcf_function_id": "attacker-via-metadata"},
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="nvidia_riva/parakeet",
|
|
)
|
|
|
|
def test_nested_langfuse_host_in_embedding_config_blocked(self):
|
|
"""The recursion uses the *full* banned-param list, not a special
|
|
subset — so any flag that's banned at the root is also banned
|
|
when nested."""
|
|
with pytest.raises(ValueError, match="langfuse_host"):
|
|
is_request_body_safe(
|
|
request_body={"litellm_embedding_config": {"langfuse_host": "https://attacker.example.com"}},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="milvus-store",
|
|
)
|
|
|
|
def test_nested_api_base_allowed_when_admin_opts_in(self):
|
|
"""Admins who explicitly enable client-side credential passthrough
|
|
keep the existing escape hatch — same UX as for root-level."""
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"litellm_embedding_config": {"api_base": "https://my-azure.example.com"}},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="milvus-store",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_safe_nested_config_accepted(self):
|
|
"""A nested config without any banned params passes — there's no
|
|
false-positive on legitimate ``api_version`` / model params."""
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"litellm_embedding_config": {
|
|
"api_version": "2024-02-15-preview",
|
|
}
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="milvus-store",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_non_dict_nested_config_does_not_break_check(self):
|
|
"""A bogus type for ``litellm_embedding_config`` (string, list,
|
|
None) must not crash the validator — it should just fall through."""
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"litellm_embedding_config": "not-a-dict"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="x",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_deeply_nested_config_does_not_recurse(self):
|
|
"""Greptile P1: ``is_request_body_safe`` is iterative single-level —
|
|
a deeply-nested ``litellm_embedding_config`` cannot exhaust the
|
|
Python call stack to trigger a 500 ``RecursionError``. Build a
|
|
body 1000 levels deep; the validator must complete in O(1)
|
|
descent."""
|
|
body = {"litellm_embedding_config": {}}
|
|
cur = body["litellm_embedding_config"]
|
|
for _ in range(1000):
|
|
cur["litellm_embedding_config"] = {}
|
|
cur = cur["litellm_embedding_config"]
|
|
# Banned param at the deepest level shouldn't be reached — single
|
|
# level only.
|
|
cur["api_base"] = "https://attacker.example.com"
|
|
|
|
# No exception raised: deeper levels aren't checked.
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body=body,
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="x",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
# ── observability-callback ban (root + metadata) ───────────────────────────
|
|
|
|
|
|
class TestObservabilityCallbackBans:
|
|
"""The proxy must reject observability credentials, hosts, and project
|
|
identifiers regardless of whether they arrive at the request body root,
|
|
in ``metadata`` / ``litellm_metadata``, or in a JSON-string-encoded
|
|
metadata blob (multipart/``extra_body`` path).
|
|
|
|
The ban list is derived from
|
|
``litellm.litellm_core_utils.initialize_dynamic_callback_params._supported_callback_params``
|
|
minus a small ``_SAFE_CLIENT_CALLBACK_PARAMS`` allow-list, plus
|
|
``_EXTRA_BANNED_OBSERVABILITY_PARAMS`` for fields integrations read but
|
|
that are not yet in the canonical allow-list. The derivation keeps the
|
|
proxy in sync as new integrations are added.
|
|
"""
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _disable_url_validation(self, monkeypatch):
|
|
import litellm
|
|
|
|
monkeypatch.setattr(litellm, "user_url_validation", False, raising=False)
|
|
|
|
@pytest.mark.parametrize(
|
|
"field",
|
|
[
|
|
"langfuse_public_key",
|
|
"langfuse_secret",
|
|
"langfuse_secret_key",
|
|
"langsmith_api_key",
|
|
"langsmith_project",
|
|
"langsmith_tenant_id",
|
|
"arize_api_key",
|
|
"arize_space_key",
|
|
"arize_space_id",
|
|
"posthog_api_key",
|
|
"posthog_api_url",
|
|
"braintrust_api_key",
|
|
"braintrust_project",
|
|
"wandb_api_key",
|
|
"weave_project_id",
|
|
"gcs_bucket_name",
|
|
"gcs_path_service_account",
|
|
"humanloop_api_key",
|
|
"lunary_public_key",
|
|
],
|
|
)
|
|
def test_observability_field_in_request_body_root_is_rejected(self, field):
|
|
with pytest.raises(ValueError, match="Rejected Request") as exc:
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", field: "attacker-value"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
assert field in str(exc.value)
|
|
|
|
@pytest.mark.parametrize(
|
|
"metadata_key",
|
|
["metadata", "litellm_metadata"],
|
|
)
|
|
@pytest.mark.parametrize(
|
|
"field",
|
|
[
|
|
"langfuse_host",
|
|
"langfuse_secret_key",
|
|
"langsmith_api_key",
|
|
"posthog_api_url",
|
|
"braintrust_project",
|
|
"user_api_key_auth_metadata",
|
|
],
|
|
)
|
|
def test_observability_field_in_metadata_dict_is_rejected(self, metadata_key, field):
|
|
# Verifies the metadata walk: a value smuggled inside ``metadata``
|
|
# or ``litellm_metadata`` is just as dangerous as the same field
|
|
# at the body root, and must hit the same gate.
|
|
with pytest.raises(ValueError, match="Rejected Request") as exc:
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
metadata_key: {field: "attacker-value"},
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
assert field in str(exc.value)
|
|
|
|
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
|
|
@pytest.mark.parametrize(
|
|
"field",
|
|
["phoenix_project_name", "phoenix_project_name_override"],
|
|
)
|
|
def test_phoenix_project_fields_in_metadata_are_accepted(self, metadata_key, field):
|
|
# The Phoenix integrations only honor the project from
|
|
# ``user_api_key_auth_metadata`` on the proxy, so the bare metadata
|
|
# fields are inert and must not 400 SDK-style callers that send them.
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
metadata_key: {field: "client-project"},
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_observability_field_in_litellm_params_metadata_is_rejected(self):
|
|
with pytest.raises(ValueError, match="Rejected Request: turn_off_message_logging is not allowed") as exc:
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"litellm_params": {"metadata": {"turn_off_message_logging": False}},
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
assert "turn_off_message_logging" in str(exc.value)
|
|
|
|
@pytest.mark.parametrize(
|
|
"metadata_key",
|
|
["metadata", "litellm_metadata"],
|
|
)
|
|
def test_observability_field_in_json_string_metadata_is_rejected(self, metadata_key):
|
|
# Multipart/form-data and ``extra_body`` callers send metadata as a
|
|
# JSON-encoded string. The bouncer parses it before applying the
|
|
# banned-params check so the JSON-string path can't smuggle past
|
|
# the ``isinstance(dict)`` guard.
|
|
import json
|
|
|
|
with pytest.raises(ValueError, match="Rejected Request: langfuse_host is not allowed in request") as exc:
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
metadata_key: json.dumps({"langfuse_host": "https://attacker.example"}),
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
assert "langfuse_host" in str(exc.value)
|
|
|
|
def test_admin_opt_in_allows_metadata_credential_passthrough(self):
|
|
# The opt-in gate covers the metadata path the same way it covers
|
|
# the root path — operators running BYO observability with
|
|
# clientside creds flip a single flag and both paths work.
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"metadata": {
|
|
"langfuse_host": "https://my-langfuse.example",
|
|
"langfuse_public_key": "pk-mine",
|
|
"langfuse_secret_key": "sk-mine",
|
|
},
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_safe_per_request_observability_metadata_is_allowed(self):
|
|
# Informational fields (sampling rate, prompt version) describe
|
|
# the request being logged — they don't choose the destination or
|
|
# credentials, so they must remain accepted from clients without
|
|
# the opt-in flag.
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"metadata": {
|
|
"langfuse_prompt_version": "v2",
|
|
"langsmith_sampling_rate": 0.1,
|
|
},
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
def test_model_level_allow_does_not_skip_subsequent_banned_params(monkeypatch):
|
|
"""Greptile P1: ``_check_banned_params`` previously ``return``-ed when a
|
|
deployment's ``configurable_clientside_auth_params`` permitted one
|
|
banned field, exiting before any later banned field in the same body
|
|
was checked. The metadata walk this PR adds multiplies the surface
|
|
where that bypass matters: a body pairing a model-level-allowed
|
|
``api_base`` with an observability credential like ``langfuse_host``
|
|
must still reject on the second field, not silently pass."""
|
|
from litellm.proxy.auth import auth_utils
|
|
|
|
monkeypatch.setattr(
|
|
auth_utils,
|
|
"_allow_model_level_clientside_configurable_parameters",
|
|
lambda model, param, request_body_value, llm_router: param == "api_base",
|
|
)
|
|
|
|
with pytest.raises(ValueError, match="Rejected Request: langfuse_host is not allowed in request") as exc:
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"api_base": "https://allowed-by-deployment.example",
|
|
"langfuse_host": "https://attacker.example",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
assert "langfuse_host" in str(exc.value)
|
|
|
|
|
|
def test_observability_ban_covers_canonical_supported_callback_params():
|
|
"""Guard test: every entry in the canonical
|
|
``_supported_callback_params`` allow-list must end up either banned by
|
|
the proxy or explicitly safe-listed. New integrations added to that
|
|
list are banned by default (the safe failure mode); flagging them as
|
|
safe is an explicit decision recorded in
|
|
``_SAFE_CLIENT_CALLBACK_PARAMS``."""
|
|
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
|
_request_blocked_callback_params,
|
|
_supported_callback_params,
|
|
)
|
|
from litellm.proxy.auth.auth_utils import (
|
|
_BANNED_REQUEST_BODY_PARAMS,
|
|
_SAFE_CLIENT_CALLBACK_PARAMS,
|
|
)
|
|
|
|
banned = set(_BANNED_REQUEST_BODY_PARAMS)
|
|
for param in _supported_callback_params:
|
|
assert param in banned or param in _SAFE_CLIENT_CALLBACK_PARAMS, (
|
|
f"{param} is in _supported_callback_params but neither banned nor "
|
|
f"safe-listed. Add it to _SAFE_CLIENT_CALLBACK_PARAMS if it is an "
|
|
f"informational per-request field; otherwise the derivation will "
|
|
f"ban it automatically."
|
|
)
|
|
for param in _request_blocked_callback_params:
|
|
assert param in banned, (
|
|
f"{param} is in _request_blocked_callback_params but is not banned at the proxy request-body boundary."
|
|
)
|
|
|
|
|
|
# ── pricing injection (global model cost registry poisoning) ──────────────────
|
|
|
|
|
|
class TestPricingInjectionBlocked:
|
|
"""Authenticated clients must not be able to mutate the global
|
|
litellm.model_cost registry by supplying pricing fields in the request
|
|
body. Any CustomPricingLiteLLMParams field (input_cost_per_token etc.)
|
|
passed to completion() is forwarded to register_model(), which overwrites
|
|
the shared global dict for ALL users on the instance.
|
|
|
|
Fix: all CustomPricingLiteLLMParams fields are in _BANNED_REQUEST_BODY_PARAMS,
|
|
so is_request_body_safe() rejects them before they reach completion().
|
|
"""
|
|
|
|
@pytest.mark.parametrize(
|
|
"field,value",
|
|
[
|
|
("input_cost_per_token", -0.01),
|
|
("output_cost_per_token", 0.0),
|
|
("input_cost_per_second", 999.0),
|
|
("output_cost_per_second", -1.0),
|
|
("cache_read_input_token_cost", 0.0),
|
|
("cache_creation_input_token_cost", -0.05),
|
|
],
|
|
)
|
|
def test_pricing_field_rejected_by_default(self, field, value):
|
|
with pytest.raises(ValueError, match="Rejected Request") as exc:
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", field: value},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
assert field in str(exc.value)
|
|
|
|
def test_all_custom_pricing_fields_are_banned(self):
|
|
from litellm.proxy.auth.auth_utils import _BANNED_REQUEST_BODY_PARAMS
|
|
from litellm.types.utils import CustomPricingLiteLLMParams
|
|
|
|
banned = set(_BANNED_REQUEST_BODY_PARAMS)
|
|
for field in CustomPricingLiteLLMParams.model_fields:
|
|
assert field in banned, (
|
|
f"CustomPricingLiteLLMParams.{field} is not in "
|
|
"_BANNED_REQUEST_BODY_PARAMS — clients can poison the global "
|
|
"model cost registry by supplying it in the request body."
|
|
)
|
|
|
|
def test_pricing_field_allowed_with_admin_opt_in(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "input_cost_per_token": 0.00001},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestGetRequestRouteTemplate:
|
|
"""get_request_route_template returns the low-cardinality FastAPI route
|
|
template (e.g. /v1/threads/{thread_id}/runs) for http.route, distinct
|
|
from the literal url.path. None when unavailable."""
|
|
|
|
def _request(self, scope):
|
|
req = MagicMock()
|
|
req.scope = scope
|
|
return req
|
|
|
|
def test_returns_route_template(self):
|
|
route = MagicMock()
|
|
route.path = "/v1/threads/{thread_id}/runs"
|
|
req = self._request({"route": route, "path": "/v1/threads/abc123/runs"})
|
|
# template, not the literal path — two thread IDs share this value
|
|
assert get_request_route_template(req) == "/v1/threads/{thread_id}/runs"
|
|
|
|
def test_scope_not_dict_returns_none(self):
|
|
assert get_request_route_template(self._request("not-a-dict")) is None
|
|
|
|
def test_no_route_in_scope_returns_none(self):
|
|
assert get_request_route_template(self._request({"path": "/x"})) is None
|
|
|
|
def test_route_without_str_path_returns_none(self):
|
|
route = MagicMock()
|
|
route.path = 12345 # not a str
|
|
assert get_request_route_template(self._request({"route": route})) is None
|
|
|
|
def test_route_with_empty_path_returns_none(self):
|
|
route = MagicMock()
|
|
route.path = ""
|
|
assert get_request_route_template(self._request({"route": route})) is None
|
|
|
|
def test_exception_returns_none(self):
|
|
req = MagicMock()
|
|
type(req).scope = property(lambda self: (_ for _ in ()).throw(RuntimeError("boom")))
|
|
assert get_request_route_template(req) is None
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksModelList:
|
|
"""model_list is an SDK-only field with no proxy API meaning; it must
|
|
be rejected from the request body regardless of any opt-in."""
|
|
|
|
def test_model_list_rejected_with_no_opt_in(self):
|
|
with pytest.raises(ValueError, match="model_list is not allowed"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"messages": [{"role": "user", "content": "hi"}],
|
|
"model_list": [{"model_name": "x", "litellm_params": {}}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_model_list_rejected_even_with_proxy_wide_opt_in(self):
|
|
with pytest.raises(ValueError, match="model_list is not allowed"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"messages": [{"role": "user", "content": "hi"}],
|
|
"model_list": [],
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_normal_body_still_passes(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "gpt-4",
|
|
"messages": [{"role": "user", "content": "hi"}],
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestGetKeyTagRateLimits:
|
|
"""Tests for get_key_tag_rpm_limit."""
|
|
|
|
def test_reads_tag_rpm_limit_from_metadata(self):
|
|
key = UserAPIKeyAuth(api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}})
|
|
assert get_key_tag_rpm_limit(key) == {"cell-1": 5}
|
|
|
|
def test_returns_none_when_unset(self):
|
|
key = UserAPIKeyAuth(api_key="sk-123")
|
|
assert get_key_tag_rpm_limit(key) is None
|
|
|
|
|
|
class TestIsRequestBodySafeChecksBracketNotationMetadata:
|
|
"""Bracket notation is how multipart callers express nested metadata; it is
|
|
validated the same way the dict form is."""
|
|
|
|
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
|
|
def test_bracket_notation_banned_param_is_rejected(self, metadata_key):
|
|
with pytest.raises(ValueError, match="langfuse_host"):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"purpose": "assistants",
|
|
f"{metadata_key}[langfuse_host]": "https://example.invalid",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_bracket_notation_api_base_is_rejected(self):
|
|
with pytest.raises(ValueError, match="api_base"):
|
|
is_request_body_safe(
|
|
request_body={"litellm_metadata[api_base]": "https://example.invalid"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
|
|
def test_bracket_notation_allowed_under_proxy_wide_opt_in(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"litellm_metadata[langfuse_host]": "https://byok.example"},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_benign_bracket_notation_metadata_is_allowed(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"purpose": "assistants",
|
|
"litellm_metadata[spend_logs_metadata][owner]": "john",
|
|
"litellm_metadata[tags]": "production",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_bracket_notation_matches_json_encoding_for_deeper_nesting(self):
|
|
"""A value nested below the first level is treated the same either way:
|
|
the check descends one level into metadata, for both encodings."""
|
|
deep_bracket = {"litellm_metadata[spend_logs_metadata][langfuse_host]": "https://example.invalid"}
|
|
deep_json = {"litellm_metadata": {"spend_logs_metadata": {"langfuse_host": "https://example.invalid"}}}
|
|
kwargs = dict(general_settings={}, llm_router=None, model="gpt-4")
|
|
assert is_request_body_safe(request_body=deep_bracket, **kwargs) is True
|
|
assert is_request_body_safe(request_body=deep_json, **kwargs) is True
|
|
|
|
def test_body_without_bracket_keys_is_unaffected(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="gpt-4",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
class TestHasUserSetupSso:
|
|
"""has_user_setup_sso must treat SAML IdP metadata as SSO configured.
|
|
|
|
Regression: UI discovery used this helper for sso_configured, but it only
|
|
checked OAuth client IDs, so SAML-only setups left the login button gray.
|
|
"""
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _clear_sso_env(self, monkeypatch):
|
|
for key in (
|
|
"MICROSOFT_CLIENT_ID",
|
|
"GOOGLE_CLIENT_ID",
|
|
"GENERIC_CLIENT_ID",
|
|
"SAML_IDP_METADATA_URL",
|
|
"SAML_IDP_METADATA_XML",
|
|
):
|
|
monkeypatch.delenv(key, raising=False)
|
|
|
|
def test_false_when_no_sso_env(self):
|
|
from litellm.proxy.auth.auth_utils import has_user_setup_sso
|
|
|
|
assert has_user_setup_sso() is False
|
|
|
|
def test_true_for_oauth_client_ids(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import has_user_setup_sso
|
|
|
|
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-client")
|
|
assert has_user_setup_sso() is True
|
|
|
|
def test_true_for_saml_metadata_url(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import has_user_setup_sso
|
|
|
|
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml")
|
|
assert has_user_setup_sso() is True
|
|
|
|
def test_true_for_saml_metadata_xml(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import has_user_setup_sso
|
|
|
|
monkeypatch.setenv("SAML_IDP_METADATA_XML", "<EntityDescriptor/>")
|
|
assert has_user_setup_sso() is True
|
|
|
|
|
|
class TestIsSsoProviderFullyConfigured:
|
|
"""A lone client id must not read as ready: `has_user_setup_sso()` only
|
|
checks the client id (correct for a UI-discovery "show the login button"
|
|
decision), but a gate that BLOCKS the password fallback needs every
|
|
companion setting the provider requires, or an incomplete setup locks
|
|
every admin out with no working login path at all."""
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _clear_sso_env(self, monkeypatch):
|
|
for key in (
|
|
"GOOGLE_CLIENT_ID",
|
|
"GOOGLE_CLIENT_SECRET",
|
|
"MICROSOFT_CLIENT_ID",
|
|
"MICROSOFT_CLIENT_SECRET",
|
|
"MICROSOFT_TENANT",
|
|
"GENERIC_CLIENT_ID",
|
|
"GENERIC_CLIENT_SECRET",
|
|
"GENERIC_AUTHORIZATION_ENDPOINT",
|
|
"GENERIC_TOKEN_ENDPOINT",
|
|
"GENERIC_USERINFO_ENDPOINT",
|
|
"SAML_IDP_METADATA_URL",
|
|
"SAML_IDP_METADATA_XML",
|
|
):
|
|
monkeypatch.delenv(key, raising=False)
|
|
|
|
def test_false_when_nothing_configured(self):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
assert is_sso_provider_fully_configured() is False
|
|
|
|
def test_google_client_id_alone_is_not_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-client")
|
|
assert is_sso_provider_fully_configured() is False
|
|
|
|
def test_google_with_secret_is_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-client")
|
|
monkeypatch.setenv("GOOGLE_CLIENT_SECRET", "google-secret")
|
|
assert is_sso_provider_fully_configured() is True
|
|
|
|
def test_microsoft_client_id_alone_is_not_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-client")
|
|
assert is_sso_provider_fully_configured() is False
|
|
|
|
def test_microsoft_missing_tenant_is_not_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-client")
|
|
monkeypatch.setenv("MICROSOFT_CLIENT_SECRET", "ms-secret")
|
|
assert is_sso_provider_fully_configured() is False
|
|
|
|
def test_microsoft_with_secret_and_tenant_is_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-client")
|
|
monkeypatch.setenv("MICROSOFT_CLIENT_SECRET", "ms-secret")
|
|
monkeypatch.setenv("MICROSOFT_TENANT", "ms-tenant")
|
|
assert is_sso_provider_fully_configured() is True
|
|
|
|
def test_generic_client_id_alone_is_not_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-client")
|
|
assert is_sso_provider_fully_configured() is False
|
|
|
|
def test_generic_missing_one_endpoint_is_not_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-client")
|
|
monkeypatch.setenv("GENERIC_CLIENT_SECRET", "generic-secret")
|
|
monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://idp.example.com/authorize")
|
|
monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://idp.example.com/token")
|
|
# GENERIC_USERINFO_ENDPOINT deliberately left unset.
|
|
assert is_sso_provider_fully_configured() is False
|
|
|
|
def test_generic_with_every_endpoint_is_ready(self, monkeypatch):
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-client")
|
|
monkeypatch.setenv("GENERIC_CLIENT_SECRET", "generic-secret")
|
|
monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://idp.example.com/authorize")
|
|
monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://idp.example.com/token")
|
|
monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://idp.example.com/userinfo")
|
|
assert is_sso_provider_fully_configured() is True
|
|
|
|
def test_saml_metadata_url_is_ready_when_runtime_installed(self, monkeypatch):
|
|
from litellm.proxy.auth import auth_utils
|
|
|
|
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml")
|
|
monkeypatch.setattr(auth_utils.importlib.util, "find_spec", lambda name: object())
|
|
assert auth_utils.is_sso_provider_fully_configured() is True
|
|
|
|
def test_saml_metadata_url_is_not_ready_without_runtime(self, monkeypatch):
|
|
"""Regression: python3-saml (``onelogin.saml2``) is an optional
|
|
dependency; SAMLAuthHandler fails closed on every request when it is
|
|
not installed, so IdP metadata alone must not read as ready."""
|
|
from litellm.proxy.auth import auth_utils
|
|
|
|
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml")
|
|
monkeypatch.setattr(auth_utils.importlib.util, "find_spec", lambda name: None)
|
|
assert auth_utils.is_sso_provider_fully_configured() is False
|
|
|
|
def test_saml_check_does_not_raise_when_package_entirely_absent(self, monkeypatch):
|
|
"""Regression: `importlib.util.find_spec("onelogin.saml2.auth")`
|
|
raises ModuleNotFoundError (not merely returns None) when the
|
|
TOP-LEVEL `onelogin` package is not installed at all, which is
|
|
exactly the real-world "optional extra not installed" case. If the
|
|
gate does not catch this, every password login 500s instead of
|
|
falling back, on a deployment that configured SAML metadata but
|
|
skipped the extra."""
|
|
from litellm.proxy.auth import auth_utils
|
|
|
|
def _raise(name: str):
|
|
raise ModuleNotFoundError("No module named 'onelogin'")
|
|
|
|
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml")
|
|
monkeypatch.setattr(auth_utils.importlib.util, "find_spec", _raise)
|
|
assert auth_utils.is_sso_provider_fully_configured() is False
|
|
|
|
def test_incomplete_earlier_provider_does_not_mask_a_ready_later_one(self, monkeypatch):
|
|
"""Regression: a stray GOOGLE_CLIENT_ID with no secret (e.g. a
|
|
leftover from a migration) must not stop the check from reaching a
|
|
fully configured Microsoft provider set alongside it — every
|
|
provider is evaluated independently, not in a first-match order."""
|
|
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
|
|
|
|
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-client")
|
|
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-client")
|
|
monkeypatch.setenv("MICROSOFT_CLIENT_SECRET", "ms-secret")
|
|
monkeypatch.setenv("MICROSOFT_TENANT", "ms-tenant")
|
|
assert is_sso_provider_fully_configured() is True
|
|
|
|
|
|
class TestIsRequestBodySafeBlocksAwsIdentitySelectors:
|
|
"""A caller must not be able to redirect Bedrock signing to another identity
|
|
reachable from the proxy host. ``get_credentials`` prefers a named profile
|
|
and the AssumeRole knobs over the deployment's static keys, and the file /
|
|
batch endpoints fold the request body and the deployment credentials into a
|
|
single params dict, so these have to be rejected at the boundary (#36155).
|
|
"""
|
|
|
|
@pytest.mark.parametrize(
|
|
"selector",
|
|
["aws_profile_name", "aws_session_name", "aws_external_id", "aws_session_tags"],
|
|
)
|
|
def test_aws_identity_selector_in_batch_body_is_rejected(self, selector):
|
|
with pytest.raises(ValueError, match=selector):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"input_file_id": "file-abc123",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"model": "bedrock-batch-model",
|
|
selector: "attacker-chosen",
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="bedrock-batch-model",
|
|
)
|
|
|
|
@pytest.mark.parametrize(
|
|
"selector",
|
|
["aws_profile_name", "aws_session_name", "aws_external_id", "aws_session_tags"],
|
|
)
|
|
def test_aws_identity_selector_under_extra_body_is_rejected(self, selector):
|
|
with pytest.raises(ValueError, match=selector):
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "bedrock-batch-model",
|
|
"extra_body": {selector: "attacker-chosen"},
|
|
},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="bedrock-batch-model",
|
|
)
|
|
|
|
def test_aws_identity_selector_allowed_under_proxy_wide_opt_in(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={
|
|
"model": "bedrock-batch-model",
|
|
"aws_profile_name": "admin-approved-profile",
|
|
},
|
|
general_settings={"allow_client_side_credentials": True},
|
|
llm_router=None,
|
|
model="bedrock-batch-model",
|
|
)
|
|
is True
|
|
)
|
|
|
|
def test_upload_body_without_identity_selectors_is_accepted(self):
|
|
assert (
|
|
is_request_body_safe(
|
|
request_body={"purpose": "batch", "model": "bedrock-batch-model"},
|
|
general_settings={},
|
|
llm_router=None,
|
|
model="bedrock-batch-model",
|
|
)
|
|
is True
|
|
)
|