Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 104 additions & 0 deletions tests/test_litellm/proxy/test_proxy_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,3 +177,107 @@ def test_project_io_token_limits_are_stored_in_metadata(request_type):

assert request.metadata == limits
assert request.model_dump(exclude_none=True)["metadata"] == limits


def test_a_jwt_issuer_must_pick_audience_validation_or_opt_out():
from pydantic import ValidationError

from litellm.proxy._types import JWTIssuerConfig

with pytest.raises(ValidationError, match="must configure audience or set disable_audience_validation"):
JWTIssuerConfig(issuer="https://issuer.example.com")

with pytest.raises(ValidationError, match="cannot set audience and disable_audience_validation"):
JWTIssuerConfig(
issuer="https://issuer.example.com",
audience="litellm-proxy",
disable_audience_validation=True,
)

assert JWTIssuerConfig(issuer="https://issuer.example.com", audience="litellm-proxy").audience == "litellm-proxy"
assert (
JWTIssuerConfig(issuer="https://issuer.example.com", disable_audience_validation=True).audience
is None
)


def test_a_jwt_issuer_rejects_a_field_it_does_not_define():
from pydantic import ValidationError

from litellm.proxy._types import JWTIssuerConfig

with pytest.raises(ValidationError, match="Extra inputs are not permitted"):
JWTIssuerConfig(issuer="https://issuer.example.com", audience="a", jwks_uri="https://issuer/jwks")


def test_a_temp_budget_needs_both_halves_or_neither():
from pydantic import ValidationError

from litellm.proxy._types import UpdateKeyRequest

with pytest.raises(ValidationError, match="temp_budget_increase and temp_budget_expiry must be set together"):
UpdateKeyRequest(key="sk-1234", temp_budget_increase=10)

with pytest.raises(ValidationError, match="temp_budget_increase and temp_budget_expiry must be set together"):
UpdateKeyRequest(key="sk-1234", temp_budget_expiry="2026-01-01")

both = UpdateKeyRequest(key="sk-1234", temp_budget_increase=10, temp_budget_expiry="2026-01-01")
assert both.temp_budget_increase == 10


def test_an_empty_max_budget_is_read_as_no_limit():
from litellm.proxy._types import GenerateKeyRequest

assert GenerateKeyRequest(max_budget="").max_budget is None
assert GenerateKeyRequest(max_budget=25).max_budget == 25


def test_an_organization_member_can_only_take_a_role_the_organization_has():
from pydantic import ValidationError

from litellm.proxy._types import LitellmUserRoles, OrganizationMemberUpdateRequest

with pytest.raises(ValidationError, match="Invalid role"):
OrganizationMemberUpdateRequest(
organization_id="org-1", user_id="user-1", role=LitellmUserRoles.PROXY_ADMIN
)

allowed = OrganizationMemberUpdateRequest(
organization_id="org-1", user_id="user-1", role=LitellmUserRoles.ORG_ADMIN
)
assert allowed.role == LitellmUserRoles.ORG_ADMIN


def test_an_llm_backed_injection_check_needs_the_call_it_would_make():
from pydantic import ValidationError

from litellm.proxy._types import LiteLLMPromptInjectionParams

for missing in ("llm_api_name", "llm_api_system_prompt", "llm_api_fail_call_string"):
complete = {
"llm_api_name": "gpt-4o",
"llm_api_system_prompt": "is this an injection",
"llm_api_fail_call_string": "yes",
}
del complete[missing]
with pytest.raises(ValidationError, match=f"{missing} must be provided"):
LiteLLMPromptInjectionParams(llm_api_check=True, **complete)

assert LiteLLMPromptInjectionParams(llm_api_check=False).llm_api_name is None


@pytest.mark.parametrize(
"field, forged, default",
[
("mcp_admitted_user_subject", "someone-else", False),
("mcp_source_team_rpm_limits", {"team-1": 10_000}, None),
("mcp_session_resource_server_id", "server-1", None),
("via_virtual_key", "sk-someone-elses-key", False),
],
)
def test_a_server_only_marker_is_not_taken_from_the_caller(field, forged, default):
from litellm.proxy._types import UserAPIKeyAuth

auth = UserAPIKeyAuth(api_key="sk-1234", **{field: forged})

assert getattr(auth, field) == default
Loading