From 91bfbe6efedf09f099417f1c8ba5c1e452e77e0f Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 01:45:57 +0000 Subject: [PATCH 01/41] fix(proxy): enforce organization boundaries in admin operations Validate org admin role against all requested organizations instead of returning on first match. Scope team list queries to the caller's permitted organizations when filtering by user_id. --- .../proxy/auth/auth_checks_organization.py | 16 +- .../management_endpoints/team_endpoints.py | 93 +-- .../test_team_endpoints.py | 701 ++++++++++-------- 3 files changed, 428 insertions(+), 382 deletions(-) diff --git a/litellm/proxy/auth/auth_checks_organization.py b/litellm/proxy/auth/auth_checks_organization.py index 50efe137209..d89afcffa9a 100644 --- a/litellm/proxy/auth/auth_checks_organization.py +++ b/litellm/proxy/auth/auth_checks_organization.py @@ -144,7 +144,7 @@ def _user_is_org_admin( user_object: Optional[LiteLLM_UserTable] = None, ) -> bool: """ - Helper function to check if user is an org admin for any of the passed organizations. + Helper function to check if user is an org admin for all of the passed organizations. Checks both: - `organization_id` (singular string) — legacy callers @@ -168,9 +168,13 @@ def _user_is_org_admin( if not candidate_org_ids: return False - for _membership in user_object.organization_memberships: - if _membership.organization_id in candidate_org_ids: - if _membership.user_role == LitellmUserRoles.ORG_ADMIN.value: - return True + # Build set of orgs where user is admin + admin_org_ids = { + _membership.organization_id + for _membership in user_object.organization_memberships + if _membership.user_role == LitellmUserRoles.ORG_ADMIN.value + and _membership.organization_id is not None + } - return False + # User must be admin of ALL requested orgs, not just any one + return all(org_id in admin_org_ids for org_id in candidate_org_ids) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 138469312e1..edd92cb83de 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -120,6 +120,7 @@ def _sanitize_for_log(value: Any) -> str: text = repr(value) return text.replace("\r", "").replace("\n", "") + async def _verify_team_access( team_obj: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth, @@ -314,10 +315,8 @@ class TeamMemberBudgetHandler: return # Batch-fetch existing memberships for this team (avoids N+1 queries) - existing_memberships = ( - await prisma_client.db.litellm_teammembership.find_many( - where={"team_id": team_id} - ) + existing_memberships = await prisma_client.db.litellm_teammembership.find_many( + where={"team_id": team_id} ) existing_user_ids = {m.user_id for m in existing_memberships} @@ -1659,12 +1658,12 @@ async def update_team( # noqa: PLR0915 updated_kv["router_settings"] = safe_dumps(updated_kv["router_settings"]) updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv) - team_row: Optional[ - LiteLLM_TeamTable - ] = await prisma_client.db.litellm_teamtable.update( - where={"team_id": data.team_id}, - data=updated_kv, - include={"litellm_model_table": True}, # type: ignore + team_row: Optional[LiteLLM_TeamTable] = ( + await prisma_client.db.litellm_teamtable.update( + where={"team_id": data.team_id}, + data=updated_kv, + include={"litellm_model_table": True}, # type: ignore + ) ) if team_row is None or team_row.team_id is None: @@ -2411,13 +2410,13 @@ async def team_member_delete( ) # Fetch keys before deletion to persist them - keys_to_delete: List[ - LiteLLM_VerificationToken - ] = await prisma_client.db.litellm_verificationtoken.find_many( - where={ - "user_id": {"in": list(user_ids_to_delete)}, - "team_id": data.team_id, - } + keys_to_delete: List[LiteLLM_VerificationToken] = ( + await prisma_client.db.litellm_verificationtoken.find_many( + where={ + "user_id": {"in": list(user_ids_to_delete)}, + "team_id": data.team_id, + } + ) ) if keys_to_delete: @@ -2801,10 +2800,10 @@ async def delete_team( team_rows: List[LiteLLM_TeamTable] = [] for team_id in data.team_ids: try: - team_row_base: Optional[ - BaseModel - ] = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} + team_row_base: Optional[BaseModel] = ( + await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) ) if team_row_base is None: raise Exception @@ -2870,10 +2869,10 @@ async def delete_team( _persist_deleted_verification_tokens, ) - keys_to_delete: List[ - LiteLLM_VerificationToken - ] = await prisma_client.db.litellm_verificationtoken.find_many( - where={"team_id": {"in": data.team_ids}} + keys_to_delete: List[LiteLLM_VerificationToken] = ( + await prisma_client.db.litellm_verificationtoken.find_many( + where={"team_id": {"in": data.team_ids}} + ) ) if keys_to_delete: @@ -3110,11 +3109,11 @@ async def team_info( ) try: - team_info: Optional[ - BaseModel - ] = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id}, - include={"object_permission": True}, + team_info: Optional[BaseModel] = ( + await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id}, + include={"object_permission": True}, + ) ) if team_info is None: raise Exception @@ -3405,6 +3404,7 @@ async def _get_org_admin_org_ids( m.organization_id for m in (caller_user.organization_memberships or []) if m.user_role == LitellmUserRoles.ORG_ADMIN.value + and m.organization_id is not None ] return org_ids if org_ids else None @@ -3439,13 +3439,8 @@ async def _build_team_list_where_conditions( if organization_id: where_conditions["organization_id"] = organization_id - elif org_admin_org_ids is not None and not user_id: - # Org admin without explicit org or user filter: scope to their orgs. - # NOTE: when user_id is provided, no org filter is applied — the - # query returns all teams the target user belongs to across all - # organisations. This matches the legacy /team/list behaviour in - # _authorize_and_filter_teams which fetches direct-membership teams - # without an org constraint. + elif org_admin_org_ids is not None: + # Org admin: always scope to their orgs, even when filtering by user_id. where_conditions["organization_id"] = {"in": org_admin_org_ids} if user_id: @@ -3815,7 +3810,7 @@ async def _authorize_and_filter_teams( Authorize the /team/list request and return filtered teams. - Proxy admins: all teams (or filtered by user_id if provided). - - Org admins: teams from their orgs + teams they are direct members of. + - Org admins: teams from their orgs (scoped to user_id if provided). - Own query (user_id matches caller): teams the user is a member of. - Others: 401. """ @@ -3843,6 +3838,7 @@ async def _authorize_and_filter_teams( m.organization_id for m in (caller_user.organization_memberships or []) if m.user_role == LitellmUserRoles.ORG_ADMIN.value + and m.organization_id is not None ] if not allowed_org_ids: allowed_org_ids = None @@ -3865,20 +3861,13 @@ async def _authorize_and_filter_teams( ) if not user_id: return list(org_teams) - # Also include teams the user is a direct member of (outside their orgs) - seen_team_ids = {team.team_id for team in org_teams} - all_teams = list(org_teams) - # Prisma doesn't support filtering JSON array fields, so we fetch by membership separately - member_teams = await prisma_client.db.litellm_teamtable.find_many( - where={"team_id": {"not_in": list(seen_team_ids)}} if seen_team_ids else {}, - include={"litellm_model_table": True}, - ) - for team in member_teams: - if team.members_with_roles and any( - m.get("user_id") == user_id for m in team.members_with_roles - ): - all_teams.append(team) - return all_teams + # Filter org teams to only those where the target user is a member + return [ + team + for team in org_teams + if team.members_with_roles + and any(m.get("user_id") == user_id for m in team.members_with_roles) + ] elif user_id: # Regular user: fetch all and filter by membership (Prisma can't filter JSON arrays) response = await prisma_client.db.litellm_teamtable.find_many( diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index bee6642dec7..9b4bd790493 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -1010,12 +1010,15 @@ async def test_validate_team_member_add_permissions_non_admin(): team.organization_id = None # Mock the helper functions to return False - with patch( - "litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin", - return_value=False, - ), patch( - "litellm.proxy.management_endpoints.team_endpoints._is_available_team", - return_value=False, + with ( + patch( + "litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin", + return_value=False, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._is_available_team", + return_value=False, + ), ): # Should raise HTTPException for non-admin with pytest.raises(HTTPException) as exc_info: @@ -1257,19 +1260,17 @@ async def test_update_team_team_member_budget_not_passed_to_db(): user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( - "litellm.proxy.proxy_server.llm_router" - ) as mock_llm_router, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.proxy_logging_obj" - ) as mock_logging, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.auth.auth_checks._cache_team_object" - ) as mock_cache_team, patch( - "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" - ) as mock_upsert_budget: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, + patch("litellm.proxy.proxy_server.llm_router") as mock_llm_router, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch("litellm.proxy.auth.auth_checks._cache_team_object") as mock_cache_team, + patch( + "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" + ) as mock_upsert_budget, + ): # Setup mock prisma client mock_existing_team = MagicMock() mock_existing_team.model_dump.return_value = { @@ -1690,19 +1691,17 @@ async def test_update_team_with_team_member_budget_duration(): user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( - "litellm.proxy.proxy_server.llm_router" - ) as mock_llm_router, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.proxy_logging_obj" - ) as mock_logging, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.auth.auth_checks._cache_team_object" - ) as mock_cache_team, patch( - "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" - ) as mock_upsert_budget: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, + patch("litellm.proxy.proxy_server.llm_router") as mock_llm_router, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch("litellm.proxy.auth.auth_checks._cache_team_object") as mock_cache_team, + patch( + "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" + ) as mock_upsert_budget, + ): mock_existing_team = MagicMock() mock_existing_team.model_dump.return_value = { "team_id": "test_team_id", @@ -1777,7 +1776,9 @@ async def test_backfill_team_member_budget_entries_creates_missing_memberships() from unittest.mock import AsyncMock, MagicMock from litellm.proxy._types import Member - from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) team_id = "team-abc" budget_id = "budget-xyz" @@ -1847,7 +1848,9 @@ async def test_backfill_team_member_budget_entries_no_op_when_all_exist(): from unittest.mock import AsyncMock, MagicMock from litellm.proxy._types import Member - from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) team_id = "team-abc" budget_id = "budget-xyz" @@ -1886,7 +1889,9 @@ async def test_backfill_team_member_budget_entries_empty_members(): """ from unittest.mock import AsyncMock, MagicMock - from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) mock_prisma = MagicMock() mock_prisma.db.litellm_teammembership.find_many = AsyncMock(return_value=[]) @@ -2092,11 +2097,14 @@ async def test_bulk_team_member_add_all_users_flag(): updated_team_memberships=[], ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.team_endpoints.team_member_add", - new_callable=AsyncMock, - return_value=mock_team_response, - ) as mock_team_member_add: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch( + "litellm.proxy.management_endpoints.team_endpoints.team_member_add", + new_callable=AsyncMock, + return_value=mock_team_response, + ) as mock_team_member_add, + ): # Mock the database find_many call mock_prisma.db.litellm_usertable.find_many = AsyncMock( return_value=mock_db_users @@ -2213,12 +2221,15 @@ async def test_list_team_v2_security_check_non_admin_user(): user_id="non_admin_user_123", ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ), patch("litellm.proxy.proxy_server.proxy_logging_obj"), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_user_object", - new_callable=AsyncMock, - return_value=None, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + return_value=None, + ), ): mock_prisma_client.return_value = MagicMock() # Mock non-None prisma client @@ -2260,12 +2271,15 @@ async def test_list_team_v2_security_check_non_admin_user_other_user(): user_id="non_admin_user_123", ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ), patch("litellm.proxy.proxy_server.proxy_logging_obj"), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_user_object", - new_callable=AsyncMock, - return_value=None, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + return_value=None, + ), ): mock_prisma_client.return_value = MagicMock() # Mock non-None prisma client @@ -2305,9 +2319,11 @@ async def test_list_team_v2_security_check_non_admin_user_own_teams(): user_id="non_admin_user_123", ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ), patch("litellm.proxy.proxy_server.proxy_logging_obj"): + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + ): # Mock prisma client and database operations mock_db = Mock() mock_prisma_client.db = mock_db @@ -2509,12 +2525,15 @@ async def test_list_team_v2_org_admin_sees_org_teams(): ], ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ), patch("litellm.proxy.proxy_server.proxy_logging_obj"), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_user_object", - new_callable=AsyncMock, - return_value=mock_user, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + return_value=mock_user, + ), ): mock_db = Mock() mock_prisma.db = mock_db @@ -2592,12 +2611,15 @@ async def test_list_team_v2_org_admin_cannot_view_other_orgs(): ], ) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ), patch("litellm.proxy.proxy_server.proxy_logging_obj"), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_user_object", - new_callable=AsyncMock, - return_value=mock_user, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + return_value=mock_user, + ), ): mock_prisma.db = Mock() @@ -2680,11 +2702,14 @@ async def test_list_team_v2_org_admin_with_user_id_returns_user_teams(): return mock_org_admin return mock_target_user - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ), patch("litellm.proxy.proxy_server.proxy_logging_obj"), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_user_object", - side_effect=mock_get_user_object, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + side_effect=mock_get_user_object, + ), ): mock_db = Mock() mock_prisma.db = mock_db @@ -2714,10 +2739,10 @@ async def test_list_team_v2_org_admin_with_user_id_returns_user_teams(): assert result["total"] == 1 - # Verify the where clause filters by user's teams, not org scope + # Verify the where clause filters by user's teams AND org scope where = mock_db.litellm_teamtable.find_many.call_args.kwargs["where"] assert where["team_id"] == {"in": ["team_X", "team_Y"]} - assert "organization_id" not in where + assert where["organization_id"] == {"in": ["org_A"]} @pytest.mark.asyncio @@ -2913,15 +2938,15 @@ async def test_new_team_max_budget_exceeds_user_max_budget(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -2982,15 +3007,15 @@ async def test_new_team_max_budget_within_user_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -3111,17 +3136,18 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object" + ) as mock_get_org, + ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -3253,17 +3279,18 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object" + ) as mock_get_org, + ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -3393,13 +3420,14 @@ async def test_new_team_standalone_validates_against_user_models(monkeypatch): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -3460,15 +3488,15 @@ async def test_new_team_standalone_validates_against_user_budget(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -3534,17 +3562,18 @@ async def test_new_team_org_scoped_budget_exceeds_org_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object" + ) as mock_get_org, + ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -3613,17 +3642,18 @@ async def test_new_team_org_scoped_models_not_in_org_models(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object" + ) as mock_get_org, + ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -3688,13 +3718,14 @@ async def test_update_team_standalone_budget_exceeds_user_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Mock existing standalone team (no organization_id) mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-123" @@ -3778,16 +3809,18 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit(): mock_org.models = ["gpt-4"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ) as mock_get_org, + ): # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-456" @@ -3852,13 +3885,14 @@ async def test_update_team_standalone_models_exceeds_user_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Mock existing standalone team (no organization_id) mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-models-123" @@ -3936,16 +3970,18 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit(): mock_org.models = ["gpt-4", "gpt-3.5-turbo"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ) as mock_get_org, + ): # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-budget-123" @@ -4044,16 +4080,18 @@ async def test_update_team_org_scoped_models_bypasses_user_limit(): mock_org.models = ["gpt-4", "gpt-3.5-turbo", "claude-3-opus"] mock_org.litellm_budget_table = None - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ) as mock_get_org, + ): # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-models-123" @@ -4145,16 +4183,18 @@ async def test_update_team_org_scoped_models_not_in_org_models(): mock_org.models = ["gpt-4", "gpt-3.5-turbo"] # claude-3-opus is NOT allowed mock_org.litellm_budget_table = None - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ) as mock_get_org, + ): # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-models-fail-123" @@ -4231,16 +4271,18 @@ async def test_update_team_org_scoped_models_with_all_proxy_models(): mock_org.models = [SpecialModelNames.all_proxy_models.value] # Allows all models mock_org.litellm_budget_table = None - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), - ) as mock_get_org: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ) as mock_get_org, + ): # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-all-proxy-models-123" @@ -4333,10 +4375,10 @@ async def test_update_team_tpm_limit_exceeds_user_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), ): # Mock existing standalone team mock_existing_team = MagicMock() @@ -4397,10 +4439,10 @@ async def test_update_team_rpm_limit_exceeds_user_limit(): dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), ): # Mock existing standalone team mock_existing_team = MagicMock() @@ -4479,15 +4521,15 @@ async def test_new_team_org_scoped_tpm_exceeds_org_limit(): mock_org.models = ["gpt-4"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ), ): mock_license.is_team_count_over_limit.return_value = False mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4555,15 +4597,15 @@ async def test_new_team_org_scoped_rpm_exceeds_org_limit(): mock_org.models = ["gpt-4"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ), ): mock_license.is_team_count_over_limit.return_value = False mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4634,20 +4676,22 @@ async def test_new_team_org_scoped_tpm_rpm_bypasses_user_limit(): mock_org.models = ["gpt-4"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), - ), patch( - "litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team", - new=AsyncMock(), + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team", + new=AsyncMock(), + ), ): mock_license.is_team_count_over_limit.return_value = False mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4736,13 +4780,14 @@ async def test_update_team_org_scoped_tpm_exceeds_org_limit(): mock_org.models = ["gpt-4"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ), ): # Mock existing org-scoped team mock_existing_team = MagicMock() @@ -4822,13 +4867,14 @@ async def test_update_team_org_scoped_rpm_exceeds_org_limit(): mock_org.models = ["gpt-4"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ), ): # Mock existing org-scoped team mock_existing_team = MagicMock() @@ -4911,15 +4957,15 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(): mock_org.models = ["gpt-4"] mock_org.litellm_budget_table = mock_budget_table - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.proxy_logging_obj" - ) as mock_logging, patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object", - new=AsyncMock(return_value=mock_org), + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + new=AsyncMock(return_value=mock_org), + ), ): # Mock existing org-scoped team mock_existing_team = MagicMock() @@ -5036,17 +5082,18 @@ async def test_update_team_guardrails_with_org_id(): "teams": [], } - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), patch( - "litellm.proxy.proxy_server.premium_user", - True, # Required for guardrails feature - ), patch( - "litellm.proxy.proxy_server.llm_router", MagicMock() + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ), + patch( + "litellm.proxy.proxy_server.premium_user", + True, # Required for guardrails feature + ), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), ): # Mock existing team - must have compatible models with organization mock_existing_team = MagicMock() @@ -5601,15 +5648,15 @@ async def test_new_team_soft_budget_validation( dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server._license_check" - ) as mock_license, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server._license_check") as mock_license, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -5799,13 +5846,14 @@ async def test_update_team_soft_budget_validation( dummy_request = MagicMock(spec=Request) - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.proxy_server.user_api_key_cache" - ) as mock_cache, patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit: + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): # Mock existing team with existing budgets mock_existing_team = MagicMock() mock_existing_team.team_id = "test-team-123" @@ -6794,14 +6842,18 @@ async def test_list_team_v1_batches_key_queries(): key3 = MagicMock() key3.team_id = "team-2" - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( - "litellm.proxy.management_endpoints.team_endpoints._authorize_and_filter_teams", - new_callable=AsyncMock, - return_value=[team1, team2], - ), patch( - "litellm.proxy.management_endpoints.team_endpoints.get_all_team_memberships", - new_callable=AsyncMock, - return_value=[], + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, + patch( + "litellm.proxy.management_endpoints.team_endpoints._authorize_and_filter_teams", + new_callable=AsyncMock, + return_value=[team1, team2], + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_all_team_memberships", + new_callable=AsyncMock, + return_value=[], + ), ): async def filtered_find_many(**kwargs): @@ -7070,16 +7122,17 @@ async def test_update_team_rejects_unauthorized_caller(): from litellm.proxy._types import UpdateTeamRequest - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( - "litellm.proxy.proxy_server.llm_router" - ), patch("litellm.proxy.proxy_server.user_api_key_cache"), patch( - "litellm.proxy.proxy_server.proxy_logging_obj" - ), patch( - "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" - ), patch( - "litellm.proxy.management_endpoints.team_endpoints._is_user_org_admin_for_team", - new_callable=AsyncMock, - return_value=False, + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, + patch("litellm.proxy.proxy_server.llm_router"), + patch("litellm.proxy.proxy_server.user_api_key_cache"), + patch("litellm.proxy.proxy_server.proxy_logging_obj"), + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.management_endpoints.team_endpoints._is_user_org_admin_for_team", + new_callable=AsyncMock, + return_value=False, + ), ): mock_existing_team = MagicMock() mock_existing_team.model_dump.return_value = { From 815a2bed1af796a7639a9e72bbb88c9e288da5e5 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 02:07:20 +0000 Subject: [PATCH 02/41] test: add regression tests for cross-org admin escalation Verify that an org admin of org-A cannot operate on org-B, and that an admin of both orgs can operate on both. --- .../proxy/auth/test_route_checks.py | 142 +++++++++++++----- 1 file changed, 106 insertions(+), 36 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index f1344a302d7..bca6b9e78d9 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -1058,8 +1058,15 @@ class TestModelsRouteExemptFromDisableLLMEndpoints: local_file = os.path.join( os.path.dirname(__file__), - "..", "..", "..", "..", "enterprise", - "litellm_enterprise", "proxy", "auth", "route_checks.py", + "..", + "..", + "..", + "..", + "enterprise", + "litellm_enterprise", + "proxy", + "auth", + "route_checks.py", ) local_file = os.path.abspath(local_file) @@ -1075,10 +1082,15 @@ class TestModelsRouteExemptFromDisableLLMEndpoints: """Test that /models is allowed even when LLM API routes are disabled""" EnterpriseRouteChecks = self._get_enterprise_route_checks() - with patch.object( - EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True - ), patch.object( - EnterpriseRouteChecks, "is_management_routes_disabled", return_value=False + with ( + patch.object( + EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True + ), + patch.object( + EnterpriseRouteChecks, + "is_management_routes_disabled", + return_value=False, + ), ): # /models should NOT raise - it's exempt EnterpriseRouteChecks.should_call_route("/models") @@ -1088,10 +1100,15 @@ class TestModelsRouteExemptFromDisableLLMEndpoints: """Test that /v1/models is allowed even when LLM API routes are disabled""" EnterpriseRouteChecks = self._get_enterprise_route_checks() - with patch.object( - EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True - ), patch.object( - EnterpriseRouteChecks, "is_management_routes_disabled", return_value=False + with ( + patch.object( + EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True + ), + patch.object( + EnterpriseRouteChecks, + "is_management_routes_disabled", + return_value=False, + ), ): # /v1/models should NOT raise - it's exempt EnterpriseRouteChecks.should_call_route("/v1/models") @@ -1101,10 +1118,15 @@ class TestModelsRouteExemptFromDisableLLMEndpoints: """Test that non-exempt LLM routes like /v1/chat/completions are still blocked""" EnterpriseRouteChecks = self._get_enterprise_route_checks() - with patch.object( - EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True - ), patch.object( - EnterpriseRouteChecks, "is_management_routes_disabled", return_value=False + with ( + patch.object( + EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True + ), + patch.object( + EnterpriseRouteChecks, + "is_management_routes_disabled", + return_value=False, + ), ): with pytest.raises(HTTPException) as exc_info: EnterpriseRouteChecks.should_call_route("/v1/chat/completions") @@ -1119,10 +1141,15 @@ class TestModelsRouteExemptFromDisableLLMEndpoints: """Test that /v1/embeddings is still blocked when LLM API routes are disabled""" EnterpriseRouteChecks = self._get_enterprise_route_checks() - with patch.object( - EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True - ), patch.object( - EnterpriseRouteChecks, "is_management_routes_disabled", return_value=False + with ( + patch.object( + EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=True + ), + patch.object( + EnterpriseRouteChecks, + "is_management_routes_disabled", + return_value=False, + ), ): with pytest.raises(HTTPException) as exc_info: EnterpriseRouteChecks.should_call_route("/v1/embeddings") @@ -1134,10 +1161,15 @@ class TestModelsRouteExemptFromDisableLLMEndpoints: """Test that /models works normally when LLM API routes are not disabled""" EnterpriseRouteChecks = self._get_enterprise_route_checks() - with patch.object( - EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=False - ), patch.object( - EnterpriseRouteChecks, "is_management_routes_disabled", return_value=False + with ( + patch.object( + EnterpriseRouteChecks, "is_llm_api_route_disabled", return_value=False + ), + patch.object( + EnterpriseRouteChecks, + "is_management_routes_disabled", + return_value=False, + ), ): # Should not raise EnterpriseRouteChecks.should_call_route("/models") @@ -1359,6 +1391,38 @@ def test_non_org_admin_with_organizations_list(): assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False +def test_org_admin_cannot_escalate_to_other_org(): + """Regression: admin of org-A requesting [org-A, org-B] must be rejected.""" + user_obj = _make_org_admin_user("org-A") + assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is False + + +def test_org_admin_of_multiple_orgs_can_operate_on_both(): + """Admin of both org-A and org-B can operate on both.""" + memberships = [ + LiteLLM_OrganizationMembershipTable( + user_id="multi-admin", + organization_id="org-A", + user_role=LitellmUserRoles.ORG_ADMIN.value, + created_at=datetime(2024, 1, 1), + updated_at=datetime(2024, 1, 1), + ), + LiteLLM_OrganizationMembershipTable( + user_id="multi-admin", + organization_id="org-B", + user_role=LitellmUserRoles.ORG_ADMIN.value, + created_at=datetime(2024, 1, 1), + updated_at=datetime(2024, 1, 1), + ), + ] + user_obj = LiteLLM_UserTable( + user_id="multi-admin", + user_role=LitellmUserRoles.INTERNAL_USER.value, + organization_memberships=memberships, + ) + assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is True + + @pytest.mark.asyncio async def test_initialize_pass_through_registers_wildcard_for_auth_subpath(): """ @@ -1389,15 +1453,19 @@ async def test_initialize_pass_through_registers_wildcard_for_auth_subpath(): original_routes = LiteLLMRoutes.openai_routes.value[:] try: - with patch( - "litellm.proxy.proxy_server.app", - MagicMock(), - ), patch( - "litellm.proxy.proxy_server.premium_user", - True, - ), patch( - "litellm.proxy.proxy_server.config_passthrough_endpoints", - None, + with ( + patch( + "litellm.proxy.proxy_server.app", + MagicMock(), + ), + patch( + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( + "litellm.proxy.proxy_server.config_passthrough_endpoints", + None, + ), ): await initialize_pass_through_endpoints([endpoint_config]) @@ -1417,7 +1485,9 @@ async def test_initialize_pass_through_registers_wildcard_for_auth_subpath(): # Removing the endpoint should clean up openai_routes # remove_endpoint_routes takes endpoint_id (UUID portion of # the route key "{id}:exact:{path}:{methods}") - registered = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() + registered = ( + InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() + ) endpoint_ids = {k.split(":")[0] for k in registered} for eid in endpoint_ids: InitPassThroughEndpointHelpers.remove_endpoint_routes(eid) @@ -1427,8 +1497,8 @@ async def test_initialize_pass_through_registers_wildcard_for_auth_subpath(): LiteLLMRoutes.openai_routes.value[:] = original_routes # Clean up any routes registered during this test to avoid # polluting the module-level _registered_pass_through_routes - registered = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() + registered = ( + InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() + ) for k in registered: - InitPassThroughEndpointHelpers.remove_endpoint_routes( - k.split(":")[0] - ) + InitPassThroughEndpointHelpers.remove_endpoint_routes(k.split(":")[0]) From 74a49b527c6db53bb8f89d83cdebec576589b62b Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 02:24:10 +0000 Subject: [PATCH 03/41] fix(proxy): read guardrail config from admin metadata, fix tag routing consistency Read guardrail control flags (disable_global_guardrails, opted_out_global_guardrails) from admin-configured key metadata instead of the request body. This ensures callers cannot override admin security policies. Fix tag-based routing to enforce strict tag checks regardless of whether the request includes tags. Fix budget limiter to use the same dynamic metadata key resolution as the tag router for consistent tag extraction. --- litellm/integrations/custom_guardrail.py | 29 ++++++++++++-------- litellm/proxy/litellm_pre_call_utils.py | 6 ++++ litellm/router_strategy/budget_limiter.py | 24 ++++++++++++---- litellm/router_strategy/tag_based_routing.py | 4 +-- 4 files changed, 42 insertions(+), 21 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 6046f1bb581..1e6a46044e0 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -257,24 +257,24 @@ class CustomGuardrail(CustomLogger): def get_disable_global_guardrail(self, data: dict) -> Optional[bool]: """ - Returns True if the global guardrail should be disabled + Returns True if the global guardrail should be disabled. + + Reads from admin-configured key/team metadata only, not from + the request body, to prevent callers from disabling guardrails. """ - if "disable_global_guardrails" in data: - return data["disable_global_guardrails"] metadata = data.get("litellm_metadata") or data.get("metadata", {}) - if "disable_global_guardrails" in metadata: - return metadata["disable_global_guardrails"] - return False + admin_metadata = metadata.get("user_api_key_metadata") or {} + return admin_metadata.get("disable_global_guardrails", False) def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]: """ Returns the list of global guardrail names the team/key has opted out of. + + Reads from admin-configured key/team metadata only. """ - if "opted_out_global_guardrails" in data: - value = data["opted_out_global_guardrails"] - return value if isinstance(value, list) else [] metadata = data.get("litellm_metadata") or data.get("metadata", {}) - value = metadata.get("opted_out_global_guardrails") + admin_metadata = metadata.get("user_api_key_metadata") or {} + value = admin_metadata.get("opted_out_global_guardrails") return value if isinstance(value, list) else [] def _is_valid_response_type(self, result: Any) -> bool: @@ -417,7 +417,9 @@ class CustomGuardrail(CustomLogger): """ requested_guardrails = self.get_guardrail_from_metadata(data) disable_global_guardrail = self.get_disable_global_guardrail(data) - opted_out_global_guardrails = self.get_opted_out_global_guardrails_from_metadata(data) + opted_out_global_guardrails = ( + self.get_opted_out_global_guardrails_from_metadata(data) + ) verbose_logger.debug( "inside should_run_guardrail for guardrail=%s event_type= %s guardrail_supported_event_hooks= %s requested_guardrails= %s self.default_on= %s", self.guardrail_name, @@ -426,7 +428,10 @@ class CustomGuardrail(CustomLogger): requested_guardrails, self.default_on, ) - if self.default_on is True and self.guardrail_name in opted_out_global_guardrails: + if ( + self.default_on is True + and self.guardrail_name in opted_out_global_guardrails + ): return False if self.default_on is True and disable_global_guardrail is not True: diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index b0adf7aa6ee..95e1cbed44e 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -977,6 +977,12 @@ async def add_litellm_data_to_request( # noqa: PLR0915 "Setting client-provided x-api-key as api_key parameter (will override deployment key)" ) + # Strip internal pipeline state from user input + for _meta_key in ("metadata", "litellm_metadata"): + _user_meta = data.get(_meta_key) + if isinstance(_user_meta, dict): + _user_meta.pop("_pipeline_managed_guardrails", None) + ########################################################## # Init - Proxy Server Request # we do this as soon as entering so we track the original request diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 64dc5fe4741..261e659644a 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -29,6 +29,9 @@ from litellm.caching.redis_cache import RedisPipelineIncrementOperation from litellm.integrations.custom_logger import CustomLogger, Span from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, +) from litellm.router_utils.cooldown_callbacks import ( _get_prometheus_logger_from_callbacks, ) @@ -100,9 +103,9 @@ class RouterBudgetLimiting(CustomLogger): self.dual_cache = dual_cache self.redis_increment_operation_queue: List[RedisPipelineIncrementOperation] = [] asyncio.create_task(self.periodic_sync_in_memory_spend_with_redis()) - self.provider_budget_config: Optional[ - GenericBudgetConfigType - ] = provider_budget_config + self.provider_budget_config: Optional[GenericBudgetConfigType] = ( + provider_budget_config + ) self.deployment_budget_config: Optional[GenericBudgetConfigType] = None self.tag_budget_config: Optional[GenericBudgetConfigType] = None self._init_provider_budgets() @@ -175,7 +178,10 @@ class RouterBudgetLimiting(CustomLogger): spend_map=spend_map, potential_deployments=potential_deployments, request_tags=_get_tags_from_request_kwargs( - request_kwargs=request_kwargs + request_kwargs=request_kwargs, + metadata_variable_name=get_metadata_variable_name_from_kwargs( + request_kwargs or {} + ), ), ) @@ -333,7 +339,10 @@ class RouterBudgetLimiting(CustomLogger): # Check tag budgets if self.tag_budget_config: request_tags = _get_tags_from_request_kwargs( - request_kwargs=request_kwargs + request_kwargs=request_kwargs, + metadata_variable_name=get_metadata_variable_name_from_kwargs( + request_kwargs or {} + ), ) for _tag in request_tags: _tag_budget_config = self._get_budget_config_for_tag(_tag) @@ -459,7 +468,10 @@ class RouterBudgetLimiting(CustomLogger): response_cost=response_cost, ) - request_tags = _get_tags_from_request_kwargs(kwargs) + request_tags = _get_tags_from_request_kwargs( + kwargs, + metadata_variable_name=get_metadata_variable_name_from_kwargs(kwargs or {}), + ) if len(request_tags) > 0: for _tag in request_tags: _tag_budget_config = self._get_budget_config_for_tag(_tag) diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 1188ce9d592..b0b154fc778 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -106,9 +106,7 @@ def _match_deployment( # the strict tag check has already failed (step 1 returned None). Allow # the regex to fire only when the deployment has NO plain tags, so we never # use regex as a backdoor around the operator's strict-tag policy. - strict_tag_check_failed = ( - not match_any and bool(deployment_tags) and bool(request_tags) - ) + strict_tag_check_failed = not match_any and bool(deployment_tags) if deployment_tag_regex and header_strings and not strict_tag_check_failed: regex_match = _is_valid_deployment_tag_regex( deployment_tag_regex, header_strings From 3cd5796fc7947d0302d27d1ba214c68d0f663007 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 02:28:23 +0000 Subject: [PATCH 04/41] refactor: extract admin metadata helper, hoist loop-invariant tag resolution Extract _get_admin_metadata() in CustomGuardrail to deduplicate metadata lookup. Hoist tag resolution above the deployment loop in budget limiter. Update stale comment in tag routing. --- litellm/integrations/custom_guardrail.py | 14 +++++---- litellm/router_strategy/budget_limiter.py | 30 +++++++++++--------- litellm/router_strategy/tag_based_routing.py | 8 +++--- 3 files changed, 29 insertions(+), 23 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 1e6a46044e0..4c294f6272e 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -255,6 +255,12 @@ class CustomGuardrail(CustomLogger): f"Event hook {event_hook} is not in the supported event hooks {supported_event_hooks}" ) + @staticmethod + def _get_admin_metadata(data: dict) -> dict: + """Return the admin-configured key/team metadata from the request data.""" + metadata = data.get("litellm_metadata") or data.get("metadata", {}) + return metadata.get("user_api_key_metadata") or {} + def get_disable_global_guardrail(self, data: dict) -> Optional[bool]: """ Returns True if the global guardrail should be disabled. @@ -262,9 +268,7 @@ class CustomGuardrail(CustomLogger): Reads from admin-configured key/team metadata only, not from the request body, to prevent callers from disabling guardrails. """ - metadata = data.get("litellm_metadata") or data.get("metadata", {}) - admin_metadata = metadata.get("user_api_key_metadata") or {} - return admin_metadata.get("disable_global_guardrails", False) + return self._get_admin_metadata(data).get("disable_global_guardrails", False) def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]: """ @@ -272,9 +276,7 @@ class CustomGuardrail(CustomLogger): Reads from admin-configured key/team metadata only. """ - metadata = data.get("litellm_metadata") or data.get("metadata", {}) - admin_metadata = metadata.get("user_api_key_metadata") or {} - value = admin_metadata.get("opted_out_global_guardrails") + value = self._get_admin_metadata(data).get("opted_out_global_guardrails") return value if isinstance(value, list) else [] def _is_valid_response_type(self, result: Any) -> bool: diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 261e659644a..be27b852478 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -310,6 +310,16 @@ class RouterBudgetLimiting(CustomLogger): deployment_configs: Dict[str, GenericBudgetInfo] = {} deployment_providers: List[Optional[str]] = [] + # Resolve tags once before the loop (loop-invariant) + _request_tags: List[str] = [] + if self.tag_budget_config: + _request_tags = _get_tags_from_request_kwargs( + request_kwargs=request_kwargs, + metadata_variable_name=get_metadata_variable_name_from_kwargs( + request_kwargs or {} + ), + ) + for deployment in healthy_deployments: # Check provider budgets if self.provider_budget_config: @@ -336,20 +346,14 @@ class RouterBudgetLimiting(CustomLogger): cache_keys.append( f"deployment_spend:{model_id}:{budget_config.budget_duration}" ) - # Check tag budgets - if self.tag_budget_config: - request_tags = _get_tags_from_request_kwargs( - request_kwargs=request_kwargs, - metadata_variable_name=get_metadata_variable_name_from_kwargs( - request_kwargs or {} - ), + + # Check tag budgets (outside loop — tags are per-request, not per-deployment) + for _tag in _request_tags: + _tag_budget_config = self._get_budget_config_for_tag(_tag) + if _tag_budget_config: + cache_keys.append( + f"tag_spend:{_tag}:{_tag_budget_config.budget_duration}" ) - for _tag in request_tags: - _tag_budget_config = self._get_budget_config_for_tag(_tag) - if _tag_budget_config: - cache_keys.append( - f"tag_spend:{_tag}:{_tag_budget_config.budget_duration}" - ) return ( cache_keys, provider_configs, diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index b0b154fc778..0163f3bbd4f 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -102,10 +102,10 @@ def _match_deployment( return {"matched_via": "tags", "matched_value": matched_value} # 2. Regex match against request headers. - # When match_any=False and the deployment has both plain tags and tag_regex, - # the strict tag check has already failed (step 1 returned None). Allow - # the regex to fire only when the deployment has NO plain tags, so we never - # use regex as a backdoor around the operator's strict-tag policy. + # When match_any=False and the deployment has plain tags, the strict tag + # check either didn't run (no request tags) or failed (step 1 returned + # None). Block the regex path so it cannot circumvent the operator's + # strict-tag policy. strict_tag_check_failed = not match_any and bool(deployment_tags) if deployment_tag_regex and header_strings and not strict_tag_check_failed: regex_match = _is_valid_deployment_tag_regex( From 34e9be1ba70f8f498a6a5c7c34b3e3663aeb763d Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 02:41:41 +0000 Subject: [PATCH 05/41] fix: merge team metadata in admin helper, remove turn_off_message_logging from dynamic params Include user_api_key_team_metadata alongside user_api_key_metadata in _get_admin_metadata() so team-level guardrail settings are respected. Key-level settings take precedence over team-level. Remove turn_off_message_logging from _supported_callback_params so it cannot be set via request metadata. Admin controls logging globally or via key/team configuration. Update tests to verify user-injected guardrail flags are ignored while admin-configured flags are respected. --- litellm/integrations/custom_guardrail.py | 7 +- .../initialize_dynamic_callback_params.py | 1 - .../integrations/test_custom_guardrail.py | 118 +++++++----------- 3 files changed, 53 insertions(+), 73 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 4c294f6272e..89431847e90 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -257,9 +257,12 @@ class CustomGuardrail(CustomLogger): @staticmethod def _get_admin_metadata(data: dict) -> dict: - """Return the admin-configured key/team metadata from the request data.""" + """Return merged admin-configured key and team metadata from the request data.""" metadata = data.get("litellm_metadata") or data.get("metadata", {}) - return metadata.get("user_api_key_metadata") or {} + team_meta = metadata.get("user_api_key_team_metadata") or {} + key_meta = metadata.get("user_api_key_metadata") or {} + # Key-level settings override team-level + return {**team_meta, **key_meta} def get_disable_global_guardrail(self, data: dict) -> Optional[bool]: """ diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index 92c97a59924..563609af1e4 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -48,7 +48,6 @@ _supported_callback_params = [ "braintrust_host", "slack_webhook_url", "lunary_public_key", - "turn_off_message_logging", ] diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 3a959d599b5..4c1ef853ab4 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -173,17 +173,16 @@ class TestCustomGuardrailShouldRunGuardrail: assert result is False def test_should_run_guardrail_with_disable_global_guardrail(self): - """Test that disable_global_guardrail disables a global guardrail when set to True""" + """Test that disable_global_guardrails only works from admin metadata""" from litellm.types.guardrails import GuardrailEventHooks - # Create a guardrail with default_on=True (global guardrail) custom_guardrail = CustomGuardrail( guardrail_name="global_guardrail", default_on=True, event_hook=GuardrailEventHooks.pre_call, ) - # Test 1: Global guardrail runs by default when default_on=True + # Test 1: Global guardrail runs by default data = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -193,7 +192,7 @@ class TestCustomGuardrailShouldRunGuardrail: ) assert result is True, "Global guardrail should run when default_on=True" - # Test 2: Global guardrail is disabled when disable_global_guardrail=True at root level + # Test 2: User-injected disable at root level is IGNORED data_with_disable_root = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -203,23 +202,10 @@ class TestCustomGuardrailShouldRunGuardrail: data=data_with_disable_root, event_type=GuardrailEventHooks.pre_call ) assert ( - result is False - ), "Global guardrail should be disabled when disable_global_guardrail=True" + result is True + ), "User-injected disable_global_guardrails should be ignored" - # Test 3: Global guardrail is disabled when disable_global_guardrail=True in litellm_metadata - data_with_disable_litellm = { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "test"}], - "litellm_metadata": {"disable_global_guardrails": True}, - } - result = custom_guardrail.should_run_guardrail( - data=data_with_disable_litellm, event_type=GuardrailEventHooks.pre_call - ) - assert ( - result is False - ), "Global guardrail should be disabled when disable_global_guardrail=True in litellm_metadata" - - # Test 4: Global guardrail is disabled when disable_global_guardrail=True in metadata + # Test 3: User-injected disable in metadata is IGNORED data_with_disable_metadata = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -228,25 +214,21 @@ class TestCustomGuardrailShouldRunGuardrail: result = custom_guardrail.should_run_guardrail( data=data_with_disable_metadata, event_type=GuardrailEventHooks.pre_call ) - assert ( - result is False - ), "Global guardrail should be disabled when disable_global_guardrail=True in metadata" + assert result is True, "User-injected metadata disable should be ignored" - # Test 5: Global guardrail runs when disable_global_guardrail=False - data_with_disable_false = { + # Test 4: Admin-configured disable via user_api_key_metadata IS respected + data_with_admin_disable = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], - "disable_global_guardrails": False, + "metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}}, } result = custom_guardrail.should_run_guardrail( - data=data_with_disable_false, event_type=GuardrailEventHooks.pre_call + data=data_with_admin_disable, event_type=GuardrailEventHooks.pre_call ) - assert ( - result is True - ), "Global guardrail should still run when disable_global_guardrail=False" + assert result is False, "Admin-configured disable should be respected" def test_should_run_guardrail_with_opted_out_global_guardrails(self): - """Test the per-guardrail opt-out list for global (default_on=True) guardrails""" + """Test that per-guardrail opt-out only works from admin metadata""" from litellm.types.guardrails import GuardrailEventHooks custom_guardrail = CustomGuardrail( @@ -255,7 +237,7 @@ class TestCustomGuardrailShouldRunGuardrail: event_hook=GuardrailEventHooks.pre_call, ) - # Test 1: guardrail in the opt-out list at root level → skipped + # Test 1: User-injected opt-out at root level is IGNORED data_root = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -265,23 +247,10 @@ class TestCustomGuardrailShouldRunGuardrail: custom_guardrail.should_run_guardrail( data=data_root, event_type=GuardrailEventHooks.pre_call ) - is False + is True ) - # Test 2: guardrail in the opt-out list inside litellm_metadata → skipped - data_litellm = { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "test"}], - "litellm_metadata": {"opted_out_global_guardrails": ["global_guardrail"]}, - } - assert ( - custom_guardrail.should_run_guardrail( - data=data_litellm, event_type=GuardrailEventHooks.pre_call - ) - is False - ) - - # Test 3: guardrail in the opt-out list inside metadata → skipped + # Test 2: User-injected opt-out in metadata is IGNORED data_metadata = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -291,7 +260,7 @@ class TestCustomGuardrailShouldRunGuardrail: custom_guardrail.should_run_guardrail( data=data_metadata, event_type=GuardrailEventHooks.pre_call ) - is False + is True ) # Test 4: a different guardrail in the opt-out list → still runs @@ -588,7 +557,9 @@ class TestGuardrailSensitiveFieldStripping: duration=1.0, ) - logged_response = request_data["metadata"]["standard_logging_guardrail_information"][0]["guardrail_response"] + logged_response = request_data["metadata"][ + "standard_logging_guardrail_information" + ][0]["guardrail_response"] assert "secret_fields" not in logged_response assert "sk-live-SHOULD-NOT-APPEAR" not in json.dumps(logged_response) @@ -599,7 +570,12 @@ class TestGuardrailSensitiveFieldStripping: guardrail.add_standard_logging_guardrail_information_to_request_data( guardrail_json_response=[ - {"result": "ok", "secret_fields": {"raw_headers": {"authorization": "Bearer sk-secret"}}}, + { + "result": "ok", + "secret_fields": { + "raw_headers": {"authorization": "Bearer sk-secret"} + }, + }, {"result": "also_ok"}, ], request_data=request_data, @@ -608,6 +584,7 @@ class TestGuardrailSensitiveFieldStripping: ) import json + serialized = json.dumps(request_data) assert "secret_fields" not in serialized assert "sk-secret" not in serialized @@ -621,21 +598,21 @@ class TestCustomGuardrailPassthroughSupport: """ Test that async_post_call_success_deployment_hook handles raw httpx.Response objects from passthrough endpoints without crashing with TypeError. - + This tests Fix #3: TypeError: TypedDict does not support instance and class checks """ import httpx custom_guardrail = CustomGuardrail() - + # Mock the async_post_call_success_hook to return None (guardrail didn't modify response) custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None) - + # Create a mock httpx.Response object (typical passthrough response) mock_response = AsyncMock(spec=httpx.Response) mock_response.status_code = 200 mock_response.text = "Mock response" - + request_data = { "guardrails": ["test_guardrail"], "user_api_key_user_id": "test_user", @@ -644,14 +621,14 @@ class TestCustomGuardrailPassthroughSupport: "user_api_key_hash": "test_hash", "user_api_key_request_route": "passthrough_route", } - + # This should not raise TypeError: TypedDict does not support instance and class checks result = await custom_guardrail.async_post_call_success_deployment_hook( request_data=request_data, response=mock_response, call_type=CallTypes.allm_passthrough_route, ) - + # When result is None, should return the original response assert result == mock_response @@ -659,53 +636,53 @@ class TestCustomGuardrailPassthroughSupport: async def test_async_post_call_success_deployment_hook_with_none_call_type(self): """ Test that async_post_call_success_deployment_hook handles None call_type gracefully. - + This ensures that even if call_type is None (before fix #1), the guardrail doesn't crash. """ custom_guardrail = CustomGuardrail() - + # Mock the async_post_call_success_hook to return None custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None) - + mock_response = AsyncMock() - + request_data = { "guardrails": ["test_guardrail"], "user_api_key_user_id": "test_user", } - + # Call with None call_type - should not crash result = await custom_guardrail.async_post_call_success_deployment_hook( request_data=request_data, response=mock_response, call_type=None, ) - + # Should return the original response when result is None assert result == mock_response def test_is_valid_response_type_with_none(self): """ Test _is_valid_response_type helper method correctly identifies None as invalid. - + This is part of Fix #3: Safely handling TypedDict types that don't support isinstance checks. """ custom_guardrail = CustomGuardrail() - + # None should be invalid assert custom_guardrail._is_valid_response_type(None) is False def test_is_valid_response_type_with_typeddict_error(self): """ Test _is_valid_response_type gracefully handles TypeError from TypedDict. - + This tests Fix #3: When isinstance() is called with TypedDict types, it raises TypeError. The method should catch this and allow the response through. """ from litellm.types.utils import ModelResponse - + custom_guardrail = CustomGuardrail() - + # Create a valid LiteLLM response object response = ModelResponse( id="test-id", @@ -714,13 +691,12 @@ class TestCustomGuardrailPassthroughSupport: model="test-model", object="chat.completion", ) - + # This should return True (it's a valid response type or TypeError is caught) result = custom_guardrail._is_valid_response_type(response) assert result is True - class TestEventTypeLogging: """Tests for event_type logging in guardrail information.""" @@ -1014,7 +990,9 @@ class TestTracingFieldsPopulation: guardrail_json_response="blocked", request_data=request_data, guardrail_status="guardrail_intervened", - tracing_detail=GuardrailTracingDetail(policy_template="EU AI Act Article 5"), + tracing_detail=GuardrailTracingDetail( + policy_template="EU AI Act Article 5" + ), ) slg_list = request_data["metadata"]["standard_logging_guardrail_information"] From 413f89892b54c44d626f7c25be0a2546dc1557b3 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 02:47:03 +0000 Subject: [PATCH 06/41] test: update dynamic callback params test for turn_off_message_logging removal Verify turn_off_message_logging is no longer extracted from request kwargs since it is now admin-only. --- .../test_initialize_dynamic_callback_params.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py index 91969a2b8e2..55f3c2ba3aa 100644 --- a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py +++ b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py @@ -82,13 +82,18 @@ def test_env_reference_in_litellm_params_metadata_raises(): def test_non_string_values_are_not_flagged(): kwargs = { "langsmith_sampling_rate": 0.5, - "turn_off_message_logging": True, } params = initialize_standard_callback_dynamic_params(kwargs) assert params.get("langsmith_sampling_rate") == 0.5 - assert params.get("turn_off_message_logging") is True + + +def test_turn_off_message_logging_not_extracted_from_request(): + """turn_off_message_logging is admin-only — must not be settable via request.""" + kwargs = {"turn_off_message_logging": True} + params = initialize_standard_callback_dynamic_params(kwargs) + assert params.get("turn_off_message_logging") is None def test_empty_kwargs_returns_empty_params(): From 9363f36481b8602e7866398bd054a11dd9842e8e Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:13:54 +0000 Subject: [PATCH 07/41] fix(proxy): add SSRF protection via resolve-and-rewrite for user-supplied URLs Add validate_url() utility that resolves DNS once, validates all IPs against private network ranges, and rewrites the URL to connect to the validated IP directly. Prevents DNS rebinding by pinning to the resolved IP. Disable follow_redirects to prevent redirect-based SSRF bypasses. Applied to all user-supplied URL entry points: - Image URL fetching in chat completions - Token counter image dimension fetching - RAG file ingestion - MCP OpenAPI spec loading --- .../prompt_templates/image_handling.py | 19 ++- litellm/litellm_core_utils/token_counter.py | 10 +- .../mcp_server/openapi_to_mcp_generator.py | 10 +- litellm/proxy/common_utils/url_utils.py | 123 ++++++++++++++++++ litellm/rag/ingestion/base_ingestion.py | 8 +- 5 files changed, 162 insertions(+), 8 deletions(-) create mode 100644 litellm/proxy/common_utils/url_utils.py diff --git a/litellm/litellm_core_utils/prompt_templates/image_handling.py b/litellm/litellm_core_utils/prompt_templates/image_handling.py index eaf78b7bcf5..c0727699167 100644 --- a/litellm/litellm_core_utils/prompt_templates/image_handling.py +++ b/litellm/litellm_core_utils/prompt_templates/image_handling.py @@ -10,6 +10,7 @@ import litellm from litellm import verbose_logger from litellm.caching.caching import InMemoryCache from litellm.constants import MAX_IMAGE_URL_DOWNLOAD_SIZE_MB +from litellm.proxy.common_utils.url_utils import SSRFError, validate_url MAX_IMGS_IN_MEMORY = 10 @@ -81,10 +82,17 @@ async def async_convert_url_to_base64(url: str) -> str: if cached_result: return cached_result + # Resolve DNS once, validate IPs, rewrite URL to validated IP + validated_url, original_host = validate_url(url) + client = litellm.module_level_aclient for _ in range(3): try: - response = await client.get(url, follow_redirects=True) + response = await client.get( + validated_url, + headers={"Host": original_host}, + follow_redirects=False, + ) return _process_image_response(response, url) except litellm.ImageFetchError: raise @@ -106,10 +114,17 @@ def convert_url_to_base64(url: str) -> str: if cached_result: return cached_result + # Resolve DNS once, validate IPs, rewrite URL to validated IP + validated_url, original_host = validate_url(url) + client = litellm.module_level_client for _ in range(3): try: - response = client.get(url, follow_redirects=True) + response = client.get( + validated_url, + headers={"Host": original_host}, + follow_redirects=False, + ) return _process_image_response(response, url) except litellm.ImageFetchError: raise diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 09c62f2eb55..e2d2a56c698 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -30,6 +30,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.default_encoding import encoding as default_encoding from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.proxy.common_utils.url_utils import validate_url from litellm.types.llms.anthropic import ( AnthropicMessagesToolResultParam, AnthropicMessagesToolUseParam, @@ -211,9 +212,14 @@ def get_image_dimensions( """ img_data = None try: - # Try to open as URL + # Try to open as URL — validate and pin to resolved IP + validated_url, original_host = validate_url(data) client = _get_httpx_client() - response = client.get(data) + response = client.get( + validated_url, + headers={"Host": original_host}, + follow_redirects=False, + ) img_data = response.read() except Exception: # If not URL, assume it's base64 diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 4b4818892bb..68f52f34395 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -15,6 +15,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy.common_utils.url_utils import validate_url from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -74,10 +75,13 @@ def load_openapi_spec(filepath: str) -> Dict[str, Any]: async def load_openapi_spec_async(filepath: str) -> Dict[str, Any]: if filepath.startswith("http://") or filepath.startswith("https://"): + validated_url, original_host = validate_url(filepath) client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) - # NOTE: do not close shared client if get_async_httpx_client returns a shared singleton. - # If it returns a new client each time, consider wrapping it in an async context manager. - r = await client.get(filepath) + r = await client.get( + validated_url, + headers={"Host": original_host}, + follow_redirects=False, + ) r.raise_for_status() return r.json() diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/proxy/common_utils/url_utils.py new file mode 100644 index 00000000000..97fee0c0965 --- /dev/null +++ b/litellm/proxy/common_utils/url_utils.py @@ -0,0 +1,123 @@ +""" +URL validation for user-controlled URLs. + +Use validate_url() before fetching any URL that originates from user +input (image_url, file_url, spec_path, etc.) to prevent SSRF attacks. + +The function resolves DNS once, validates all IPs, and rewrites the URL +to connect to the validated IP directly — no TOCTOU gap, no DNS rebinding. +Callers should also set follow_redirects=False to prevent redirect-based +SSRF bypasses. +""" + +import ipaddress +import socket +from ipaddress import ip_address, ip_network +from typing import Optional, Tuple +from urllib.parse import urlparse, urlunparse + +_BLOCKED_NETWORKS = [ + ip_network("0.0.0.0/8"), + ip_network("10.0.0.0/8"), + ip_network("100.64.0.0/10"), + ip_network("127.0.0.0/8"), + ip_network("169.254.0.0/16"), + ip_network("172.16.0.0/12"), + ip_network("192.0.0.0/24"), + ip_network("192.168.0.0/16"), + ip_network("198.18.0.0/15"), + ip_network("::1/128"), + ip_network("fc00::/7"), + ip_network("fe80::/10"), +] + +_ALLOWED_SCHEMES = ("http", "https") + + +class SSRFError(ValueError): + """Raised when a URL targets a blocked network.""" + + pass + + +def _is_blocked_ip(addr: str) -> bool: + try: + ip = ip_address(addr) + except ValueError: + return False + if ip.version == 6 and hasattr(ip, "ipv4_mapped") and ip.ipv4_mapped: + ip = ip.ipv4_mapped + return any(ip in net for net in _BLOCKED_NETWORKS) + + +def validate_url(url: str) -> Tuple[str, str]: + """ + Validate a user-supplied URL and rewrite it to connect to a validated IP. + + Resolves the hostname, checks all resolved IPs against blocked networks, + then returns a rewritten URL that points to the validated IP along with + the original hostname (for use in the Host header). + + This eliminates DNS rebinding because the caller connects to the IP we + validated, not the hostname that could rebind. Callers should also disable + follow_redirects to prevent redirect-based SSRF bypasses. + + Args: + url: The user-supplied URL to validate. + + Returns: + Tuple of (rewritten_url, original_hostname). + The rewritten URL has the hostname replaced with the validated IP. + The original hostname should be set as the Host header. + + Raises: + SSRFError: If the URL scheme is invalid or the hostname resolves + to a private/internal IP address. + """ + parsed = urlparse(url) + + if parsed.scheme not in _ALLOWED_SCHEMES: + raise SSRFError(f"URL scheme '{parsed.scheme}' is not allowed") + + hostname = parsed.hostname + if not hostname: + raise SSRFError("URL has no hostname") + + port = parsed.port + default_port = 443 if parsed.scheme == "https" else 80 + + # Resolve hostname and validate ALL addresses + try: + addrinfo = socket.getaddrinfo( + hostname, port or default_port, proto=socket.IPPROTO_TCP + ) + except socket.gaierror as e: + raise SSRFError(f"DNS resolution failed for '{hostname}': {e}") + + if not addrinfo: + raise SSRFError(f"No addresses found for '{hostname}'") + + for family, type_, proto, canonname, sockaddr in addrinfo: + if _is_blocked_ip(sockaddr[0]): + raise SSRFError( + f"URL targets a blocked address ({sockaddr[0]}). " + "If this is a legitimate internal service, use a direct " + "provider configuration instead of a user-supplied URL." + ) + + # Rewrite URL to connect to the first validated IP + validated_ip = addrinfo[0][4][0] + is_ipv6 = addrinfo[0][0] == socket.AF_INET6 + ip_host = f"[{validated_ip}]" if is_ipv6 else validated_ip + + # Reconstruct netloc with IP instead of hostname + if port: + new_netloc = f"{ip_host}:{port}" + else: + new_netloc = ip_host + + rewritten = urlunparse( + (parsed.scheme, new_netloc, parsed.path, parsed.params, parsed.query, "") + ) + + return rewritten, hostname diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index 0d12bdfffc1..1d868d35c7a 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -24,6 +24,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy.common_utils.url_utils import validate_url from litellm.rag.ingestion.file_parsers import extract_text_from_pdf from litellm.rag.text_splitters import RecursiveCharacterTextSplitter from litellm.types.rag import RAGIngestOptions, RAGIngestResponse @@ -111,8 +112,13 @@ class BaseRAGIngestion(ABC): return filename, file_content, content_type, None if file_url: + validated_url, original_host = validate_url(file_url) http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.RAG) - response = await http_client.get(file_url) + response = await http_client.get( + validated_url, + headers={"Host": original_host}, + follow_redirects=False, + ) response.raise_for_status() file_content = response.content filename = file_url.split("/")[-1] or "document" From d15196b5197d89270469faabaece0dd26cf1f4b6 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:30:14 +0000 Subject: [PATCH 08/41] fix(proxy): add safe_get/async_safe_get with redirect validation Add safe_get() and async_safe_get() helpers that validate each redirect hop before following. For HTTPS, rely on TLS certificate binding instead of URL rewriting. Simplify call sites to use the new helpers. --- .../prompt_templates/image_handling.py | 20 +----- litellm/litellm_core_utils/token_counter.py | 11 +-- .../mcp_server/openapi_to_mcp_generator.py | 9 +-- litellm/proxy/common_utils/url_utils.py | 70 +++++++++++++++++-- litellm/rag/ingestion/base_ingestion.py | 9 +-- 5 files changed, 73 insertions(+), 46 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/image_handling.py b/litellm/litellm_core_utils/prompt_templates/image_handling.py index c0727699167..c036ec23ddb 100644 --- a/litellm/litellm_core_utils/prompt_templates/image_handling.py +++ b/litellm/litellm_core_utils/prompt_templates/image_handling.py @@ -10,7 +10,7 @@ import litellm from litellm import verbose_logger from litellm.caching.caching import InMemoryCache from litellm.constants import MAX_IMAGE_URL_DOWNLOAD_SIZE_MB -from litellm.proxy.common_utils.url_utils import SSRFError, validate_url +from litellm.proxy.common_utils.url_utils import async_safe_get, safe_get MAX_IMGS_IN_MEMORY = 10 @@ -82,17 +82,10 @@ async def async_convert_url_to_base64(url: str) -> str: if cached_result: return cached_result - # Resolve DNS once, validate IPs, rewrite URL to validated IP - validated_url, original_host = validate_url(url) - client = litellm.module_level_aclient for _ in range(3): try: - response = await client.get( - validated_url, - headers={"Host": original_host}, - follow_redirects=False, - ) + response = await async_safe_get(client, url) return _process_image_response(response, url) except litellm.ImageFetchError: raise @@ -114,17 +107,10 @@ def convert_url_to_base64(url: str) -> str: if cached_result: return cached_result - # Resolve DNS once, validate IPs, rewrite URL to validated IP - validated_url, original_host = validate_url(url) - client = litellm.module_level_client for _ in range(3): try: - response = client.get( - validated_url, - headers={"Host": original_host}, - follow_redirects=False, - ) + response = safe_get(client, url) return _process_image_response(response, url) except litellm.ImageFetchError: raise diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index e2d2a56c698..ad2691156a5 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -30,7 +30,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.default_encoding import encoding as default_encoding from litellm.llms.custom_httpx.http_handler import _get_httpx_client -from litellm.proxy.common_utils.url_utils import validate_url +from litellm.proxy.common_utils.url_utils import safe_get from litellm.types.llms.anthropic import ( AnthropicMessagesToolResultParam, AnthropicMessagesToolUseParam, @@ -212,14 +212,9 @@ def get_image_dimensions( """ img_data = None try: - # Try to open as URL — validate and pin to resolved IP - validated_url, original_host = validate_url(data) + # Try to open as URL with SSRF protection client = _get_httpx_client() - response = client.get( - validated_url, - headers={"Host": original_host}, - follow_redirects=False, - ) + response = safe_get(client, data) img_data = response.read() except Exception: # If not URL, assume it's base64 diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 68f52f34395..d6b2cb86b26 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -15,7 +15,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy.common_utils.url_utils import validate_url +from litellm.proxy.common_utils.url_utils import async_safe_get from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -75,13 +75,8 @@ def load_openapi_spec(filepath: str) -> Dict[str, Any]: async def load_openapi_spec_async(filepath: str) -> Dict[str, Any]: if filepath.startswith("http://") or filepath.startswith("https://"): - validated_url, original_host = validate_url(filepath) client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) - r = await client.get( - validated_url, - headers={"Host": original_host}, - follow_redirects=False, - ) + r = await async_safe_get(client, filepath) r.raise_for_status() return r.json() diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/proxy/common_utils/url_utils.py index 97fee0c0965..0583cc3a42b 100644 --- a/litellm/proxy/common_utils/url_utils.py +++ b/litellm/proxy/common_utils/url_utils.py @@ -4,16 +4,15 @@ URL validation for user-controlled URLs. Use validate_url() before fetching any URL that originates from user input (image_url, file_url, spec_path, etc.) to prevent SSRF attacks. -The function resolves DNS once, validates all IPs, and rewrites the URL -to connect to the validated IP directly — no TOCTOU gap, no DNS rebinding. -Callers should also set follow_redirects=False to prevent redirect-based -SSRF bypasses. +validate_url() resolves DNS once, validates all IPs, and rewrites the +URL to connect to the validated IP directly — no TOCTOU gap, no DNS +rebinding. Redirects are followed manually with validation at each hop. """ import ipaddress import socket from ipaddress import ip_address, ip_network -from typing import Optional, Tuple +from typing import Any, Optional, Tuple, Union from urllib.parse import urlparse, urlunparse _BLOCKED_NETWORKS = [ @@ -105,12 +104,18 @@ def validate_url(url: str) -> Tuple[str, str]: "provider configuration instead of a user-supplied URL." ) - # Rewrite URL to connect to the first validated IP + # For HTTPS, TLS certificate validation binds the connection to the + # hostname — DNS rebinding can't redirect to a different server because + # the cert wouldn't match. Return the original URL. + if parsed.scheme == "https": + return url, hostname + + # For HTTP, rewrite URL to connect to the validated IP directly + # to prevent DNS rebinding (no TLS to bind the connection). validated_ip = addrinfo[0][4][0] is_ipv6 = addrinfo[0][0] == socket.AF_INET6 ip_host = f"[{validated_ip}]" if is_ipv6 else validated_ip - # Reconstruct netloc with IP instead of hostname if port: new_netloc = f"{ip_host}:{port}" else: @@ -121,3 +126,54 @@ def validate_url(url: str) -> Tuple[str, str]: ) return rewritten, hostname + + +_MAX_REDIRECTS = 10 + + +def safe_get(client: Any, url: str, **kwargs: Any) -> Any: + """ + Fetch a user-supplied URL with SSRF protection on every redirect hop. + + Validates the initial URL and each redirect target before making the + request. No DNS rebinding (resolve-and-rewrite). No redirect bypass + (each hop validated). No breaking change for legitimate CDN redirects. + + Args: + client: An httpx.Client or httpx.AsyncClient (sync version). + url: The user-supplied URL. + **kwargs: Additional kwargs passed to client.get(). + + Returns: + The final httpx.Response. + """ + kwargs.pop("follow_redirects", None) + for _ in range(_MAX_REDIRECTS): + validated_url, original_host = validate_url(url) + response = client.get( + validated_url, + headers={**kwargs.pop("headers", {}), "Host": original_host}, + follow_redirects=False, + **kwargs, + ) + if not response.is_redirect or response.next_request is None: + return response + url = str(response.next_request.url) + raise SSRFError("Too many redirects") + + +async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any: + """Async version of safe_get.""" + kwargs.pop("follow_redirects", None) + for _ in range(_MAX_REDIRECTS): + validated_url, original_host = validate_url(url) + response = await client.get( + validated_url, + headers={**kwargs.pop("headers", {}), "Host": original_host}, + follow_redirects=False, + **kwargs, + ) + if not response.is_redirect or response.next_request is None: + return response + url = str(response.next_request.url) + raise SSRFError("Too many redirects") diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index 1d868d35c7a..6f139764a85 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy.common_utils.url_utils import validate_url +from litellm.proxy.common_utils.url_utils import async_safe_get from litellm.rag.ingestion.file_parsers import extract_text_from_pdf from litellm.rag.text_splitters import RecursiveCharacterTextSplitter from litellm.types.rag import RAGIngestOptions, RAGIngestResponse @@ -112,13 +112,8 @@ class BaseRAGIngestion(ABC): return filename, file_content, content_type, None if file_url: - validated_url, original_host = validate_url(file_url) http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.RAG) - response = await http_client.get( - validated_url, - headers={"Host": original_host}, - follow_redirects=False, - ) + response = await async_safe_get(http_client, file_url) response.raise_for_status() file_content = response.content filename = file_url.split("/")[-1] or "document" From 037fb573f77e22aa67733ea3b755d515ebcc5904 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:38:29 +0000 Subject: [PATCH 09/41] fix: preserve caller headers across redirect hops in safe_get --- litellm/proxy/common_utils/url_utils.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/proxy/common_utils/url_utils.py index 0583cc3a42b..2a160d3cec7 100644 --- a/litellm/proxy/common_utils/url_utils.py +++ b/litellm/proxy/common_utils/url_utils.py @@ -148,11 +148,12 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: The final httpx.Response. """ kwargs.pop("follow_redirects", None) + caller_headers = kwargs.pop("headers", {}) for _ in range(_MAX_REDIRECTS): validated_url, original_host = validate_url(url) response = client.get( validated_url, - headers={**kwargs.pop("headers", {}), "Host": original_host}, + headers={**caller_headers, "Host": original_host}, follow_redirects=False, **kwargs, ) @@ -165,11 +166,12 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any: """Async version of safe_get.""" kwargs.pop("follow_redirects", None) + caller_headers = kwargs.pop("headers", {}) for _ in range(_MAX_REDIRECTS): validated_url, original_host = validate_url(url) response = await client.get( validated_url, - headers={**kwargs.pop("headers", {}), "Host": original_host}, + headers={**caller_headers, "Host": original_host}, follow_redirects=False, **kwargs, ) From b94aaa72b07ce1b634b3458af620a4554fb90991 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:40:54 +0000 Subject: [PATCH 10/41] fix: skip DNS resolution for base64 data in token counter, add unit tests Check URL scheme before calling safe_get in token counter to avoid unnecessary DNS resolution on base64-encoded image data. Add 14 unit tests for validate_url covering blocked networks, scheme validation, URL rewriting, and DNS failure handling. --- litellm/litellm_core_utils/token_counter.py | 16 +++-- .../proxy/common_utils/test_url_utils.py | 64 +++++++++++++++++++ 2 files changed, 73 insertions(+), 7 deletions(-) create mode 100644 tests/test_litellm/proxy/common_utils/test_url_utils.py diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index ad2691156a5..6245828a381 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -211,13 +211,15 @@ def get_image_dimensions( Tuple[int, int]: The width and height of the image. """ img_data = None - try: - # Try to open as URL with SSRF protection - client = _get_httpx_client() - response = safe_get(client, data) - img_data = response.read() - except Exception: - # If not URL, assume it's base64 + if data.startswith(("http://", "https://")): + try: + client = _get_httpx_client() + response = safe_get(client, data) + img_data = response.read() + except Exception: + pass + if img_data is None: + # Not a URL or fetch failed — assume base64 _header, encoded = data.split(",", 1) img_data = base64.b64decode(encoded) diff --git a/tests/test_litellm/proxy/common_utils/test_url_utils.py b/tests/test_litellm/proxy/common_utils/test_url_utils.py new file mode 100644 index 00000000000..73f465cfbc4 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_url_utils.py @@ -0,0 +1,64 @@ +import pytest + +from litellm.proxy.common_utils.url_utils import SSRFError, validate_url + + +class TestValidateUrl: + def test_blocks_loopback(self): + with pytest.raises(SSRFError): + validate_url("http://127.0.0.1/test") + + def test_blocks_imds(self): + with pytest.raises(SSRFError): + validate_url("http://169.254.169.254/latest/meta-data/") + + def test_blocks_rfc1918_class_a(self): + with pytest.raises(SSRFError): + validate_url("http://10.0.1.5:8080/v1/completions") + + def test_blocks_rfc1918_class_b(self): + with pytest.raises(SSRFError): + validate_url("http://172.16.0.1/") + + def test_blocks_rfc1918_class_c(self): + with pytest.raises(SSRFError): + validate_url("http://192.168.1.1/") + + def test_blocks_file_scheme(self): + with pytest.raises(SSRFError): + validate_url("file:///etc/passwd") + + def test_blocks_ftp_scheme(self): + with pytest.raises(SSRFError): + validate_url("ftp://internal.host/data") + + def test_blocks_no_hostname(self): + with pytest.raises(SSRFError): + validate_url("http:///path") + + def test_allows_public_https(self): + rewritten, host = validate_url("https://example.com/image.png") + assert host == "example.com" + assert rewritten == "https://example.com/image.png" + + def test_rewrites_public_http_to_ip(self): + rewritten, host = validate_url("http://example.com/image.png") + assert host == "example.com" + assert "example.com" not in rewritten + + def test_preserves_path_and_query(self): + rewritten, host = validate_url("http://example.com/path?key=value") + assert "/path" in rewritten + assert "key=value" in rewritten + + def test_dns_failure_raises(self): + with pytest.raises(SSRFError, match="DNS resolution failed"): + validate_url("http://this-domain-does-not-exist-xyz123.invalid/test") + + def test_blocks_localhost_hostname(self): + with pytest.raises(SSRFError): + validate_url("http://localhost/") + + def test_blocks_ipv6_loopback(self): + with pytest.raises(SSRFError): + validate_url("http://[::1]/") From 62ec39677586ae0588e5dd919d6186e4da010659 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:44:10 +0000 Subject: [PATCH 11/41] test: mock SSRF validation in openapi spec URL test --- tests/mcp_tests/test_openapi_spec_path_url.py | 14 +++++++++++--- 1 file changed, 11 insertions(+), 3 deletions(-) diff --git a/tests/mcp_tests/test_openapi_spec_path_url.py b/tests/mcp_tests/test_openapi_spec_path_url.py index 03e9db94967..17a0022046e 100644 --- a/tests/mcp_tests/test_openapi_spec_path_url.py +++ b/tests/mcp_tests/test_openapi_spec_path_url.py @@ -55,6 +55,11 @@ def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) -> # Ensure shared/custom client path is used monkeypatch.setattr(gen, "get_async_httpx_client", fake_get_async_httpx_client) + # Bypass SSRF validation in test (example.local doesn't resolve) + monkeypatch.setattr( + gen, "async_safe_get", lambda client, url, **kw: client.get(url) + ) + # Fail loudly if someone reintroduces direct httpx.get() def boom(*args, **kwargs): raise AssertionError("Direct httpx.get() must not be used for URL spec loading") @@ -68,7 +73,9 @@ def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) -> assert handler_holder["handler"].calls == 1 -def test_load_openapi_spec_supports_local_file_path(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None: +def test_load_openapi_spec_supports_local_file_path( + tmp_path, monkeypatch: pytest.MonkeyPatch +) -> None: expected: Dict[str, Any] = { "openapi": "3.0.0", "info": {"title": "Local API", "version": "1.0.0"}, @@ -83,10 +90,11 @@ def test_load_openapi_spec_supports_local_file_path(tmp_path, monkeypatch: pytes # For local files, shared client must NOT be used. def boom_client(*args, **kwargs): - raise AssertionError("get_async_httpx_client() must not be called for local file paths") + raise AssertionError( + "get_async_httpx_client() must not be called for local file paths" + ) monkeypatch.setattr(gen, "get_async_httpx_client", boom_client) spec = gen.load_openapi_spec(str(p)) assert spec == expected - From 814d03d1cee07128411c4492a67eb0813c9eb195 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:48:12 +0000 Subject: [PATCH 12/41] fix: fail-closed on unparseable IPs, rewrite HTTPS when SSL verify disabled _is_blocked_ip now returns True (blocked) for unparseable addresses instead of False (allowed). HTTPS URLs are rewritten to validated IPs when ssl_verify is disabled, closing the DNS rebinding window that exists without TLS certificate binding. --- litellm/proxy/common_utils/url_utils.py | 15 +++++++---- .../proxy/common_utils/test_url_utils.py | 25 ++++++++++++++++++- 2 files changed, 34 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/proxy/common_utils/url_utils.py index 2a160d3cec7..72cbc776589 100644 --- a/litellm/proxy/common_utils/url_utils.py +++ b/litellm/proxy/common_utils/url_utils.py @@ -15,6 +15,8 @@ from ipaddress import ip_address, ip_network from typing import Any, Optional, Tuple, Union from urllib.parse import urlparse, urlunparse +import litellm + _BLOCKED_NETWORKS = [ ip_network("0.0.0.0/8"), ip_network("10.0.0.0/8"), @@ -43,7 +45,7 @@ def _is_blocked_ip(addr: str) -> bool: try: ip = ip_address(addr) except ValueError: - return False + return True # fail-closed: unparseable addresses are blocked if ip.version == 6 and hasattr(ip, "ipv4_mapped") and ip.ipv4_mapped: ip = ip.ipv4_mapped return any(ip in net for net in _BLOCKED_NETWORKS) @@ -104,10 +106,13 @@ def validate_url(url: str) -> Tuple[str, str]: "provider configuration instead of a user-supplied URL." ) - # For HTTPS, TLS certificate validation binds the connection to the - # hostname — DNS rebinding can't redirect to a different server because - # the cert wouldn't match. Return the original URL. - if parsed.scheme == "https": + # For HTTPS with SSL verification enabled, TLS certificate validation + # binds the connection to the hostname — DNS rebinding can't redirect + # to a different server because the cert wouldn't match. + # When SSL verification is disabled, this defense doesn't apply, so + # we rewrite to the validated IP like HTTP. + ssl_verify = getattr(litellm, "ssl_verify", True) + if parsed.scheme == "https" and ssl_verify is not False: return url, hostname # For HTTP, rewrite URL to connect to the validated IP directly diff --git a/tests/test_litellm/proxy/common_utils/test_url_utils.py b/tests/test_litellm/proxy/common_utils/test_url_utils.py index 73f465cfbc4..4dbda6a815e 100644 --- a/tests/test_litellm/proxy/common_utils/test_url_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_url_utils.py @@ -1,6 +1,18 @@ import pytest -from litellm.proxy.common_utils.url_utils import SSRFError, validate_url +import litellm +from litellm.proxy.common_utils.url_utils import SSRFError, _is_blocked_ip, validate_url + + +class TestIsBlockedIp: + def test_blocks_private(self): + assert _is_blocked_ip("10.0.0.1") is True + + def test_allows_public(self): + assert _is_blocked_ip("8.8.8.8") is False + + def test_unparseable_is_blocked(self): + assert _is_blocked_ip("not-an-ip") is True class TestValidateUrl: @@ -62,3 +74,14 @@ class TestValidateUrl: def test_blocks_ipv6_loopback(self): with pytest.raises(SSRFError): validate_url("http://[::1]/") + + def test_https_rewrites_when_ssl_verify_disabled(self, monkeypatch): + monkeypatch.setattr(litellm, "ssl_verify", False) + rewritten, host = validate_url("https://example.com/image.png") + assert host == "example.com" + assert "example.com" not in rewritten # rewritten to IP + + def test_https_not_rewritten_when_ssl_verify_enabled(self, monkeypatch): + monkeypatch.setattr(litellm, "ssl_verify", True) + rewritten, host = validate_url("https://example.com/image.png") + assert rewritten == "https://example.com/image.png" From e2a0c96663548d36ebca5a530e9140da530eb19e Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:53:09 +0000 Subject: [PATCH 13/41] fix: redirect loop was dead code, clean up imports Read Location header directly instead of response.next_request (which is None when follow_redirects=False). Resolve relative redirect URLs with httpx.URL.join(). Remove unused imports. --- litellm/proxy/common_utils/url_utils.py | 25 ++++++++++++++++++------- 1 file changed, 18 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/proxy/common_utils/url_utils.py index 72cbc776589..f18378f9007 100644 --- a/litellm/proxy/common_utils/url_utils.py +++ b/litellm/proxy/common_utils/url_utils.py @@ -9,10 +9,10 @@ URL to connect to the validated IP directly — no TOCTOU gap, no DNS rebinding. Redirects are followed manually with validation at each hop. """ -import ipaddress +import asyncio import socket from ipaddress import ip_address, ip_network -from typing import Any, Optional, Tuple, Union +from typing import Any, Tuple from urllib.parse import urlparse, urlunparse import litellm @@ -136,6 +136,17 @@ def validate_url(url: str) -> Tuple[str, str]: _MAX_REDIRECTS = 10 +def _extract_redirect_url(response: Any, request_url: str) -> str: + """Extract and resolve the redirect target from a response's Location header.""" + import httpx + + location = response.headers.get("location") + if not location: + raise SSRFError("Redirect response has no Location header") + # Resolve relative URLs against the request URL + return str(httpx.URL(request_url).join(location)) + + def safe_get(client: Any, url: str, **kwargs: Any) -> Any: """ Fetch a user-supplied URL with SSRF protection on every redirect hop. @@ -145,7 +156,7 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: (each hop validated). No breaking change for legitimate CDN redirects. Args: - client: An httpx.Client or httpx.AsyncClient (sync version). + client: An httpx.Client (sync). url: The user-supplied URL. **kwargs: Additional kwargs passed to client.get(). @@ -162,9 +173,9 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: follow_redirects=False, **kwargs, ) - if not response.is_redirect or response.next_request is None: + if not response.is_redirect: return response - url = str(response.next_request.url) + url = _extract_redirect_url(response, validated_url) raise SSRFError("Too many redirects") @@ -180,7 +191,7 @@ async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any: follow_redirects=False, **kwargs, ) - if not response.is_redirect or response.next_request is None: + if not response.is_redirect: return response - url = str(response.next_request.url) + url = _extract_redirect_url(response, validated_url) raise SSRFError("Too many redirects") From 00b25d6ca4f55890cc1fb07484c8c37205677bb1 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:55:56 +0000 Subject: [PATCH 14/41] fix: sync redirect bypass, Host header port, redirect loop dead code MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Pass follow_redirects through in HTTPHandler.get() — previously the parameter was accepted but never forwarded to the underlying httpx client, making sync redirect protection ineffective. Include port in Host header when non-default (e.g. example.com:8080). Fix redirect loop to read Location header directly instead of response.next_request (which is None when follow_redirects=False). --- litellm/llms/custom_httpx/http_handler.py | 1 + litellm/proxy/common_utils/url_utils.py | 9 +++++++-- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 489a56daf8e..03d2af72329 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1019,6 +1019,7 @@ class HTTPHandler: url, params=params, headers=headers, + follow_redirects=_follow_redirects, ) return response diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/proxy/common_utils/url_utils.py index f18378f9007..a47d7ae2ca4 100644 --- a/litellm/proxy/common_utils/url_utils.py +++ b/litellm/proxy/common_utils/url_utils.py @@ -87,6 +87,11 @@ def validate_url(url: str) -> Tuple[str, str]: port = parsed.port default_port = 443 if parsed.scheme == "https" else 80 + # Build the Host header value — include port when non-default + host_header = ( + hostname if (port is None or port == default_port) else f"{hostname}:{port}" + ) + # Resolve hostname and validate ALL addresses try: addrinfo = socket.getaddrinfo( @@ -113,7 +118,7 @@ def validate_url(url: str) -> Tuple[str, str]: # we rewrite to the validated IP like HTTP. ssl_verify = getattr(litellm, "ssl_verify", True) if parsed.scheme == "https" and ssl_verify is not False: - return url, hostname + return url, host_header # For HTTP, rewrite URL to connect to the validated IP directly # to prevent DNS rebinding (no TLS to bind the connection). @@ -130,7 +135,7 @@ def validate_url(url: str) -> Tuple[str, str]: (parsed.scheme, new_netloc, parsed.path, parsed.params, parsed.query, "") ) - return rewritten, hostname + return rewritten, host_header _MAX_REDIRECTS = 10 From 1ba2be77aed17a35ebf74d963f67040ec5e03477 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:59:25 +0000 Subject: [PATCH 15/41] refactor: move url_utils to litellm_core_utils to avoid proxy dependency SDK core modules (image_handling, token_counter) should not import from litellm.proxy. Move url_utils.py to litellm_core_utils/ so bare SDK installs without proxy dependencies still work. --- litellm/litellm_core_utils/prompt_templates/image_handling.py | 2 +- litellm/litellm_core_utils/token_counter.py | 2 +- litellm/{proxy/common_utils => litellm_core_utils}/url_utils.py | 0 .../proxy/_experimental/mcp_server/openapi_to_mcp_generator.py | 2 +- litellm/proxy/_experimental/out/404/index.html | 1 + litellm/proxy/_experimental/out/_not-found/index.html | 1 + litellm/proxy/_experimental/out/api-reference/index.html | 1 + litellm/proxy/_experimental/out/chat/index.html | 1 + .../_experimental/out/experimental/api-playground/index.html | 1 + litellm/proxy/_experimental/out/experimental/budgets/index.html | 1 + litellm/proxy/_experimental/out/experimental/caching/index.html | 1 + .../out/experimental/claude-code-plugins/index.html | 1 + .../proxy/_experimental/out/experimental/old-usage/index.html | 1 + litellm/proxy/_experimental/out/experimental/prompts/index.html | 1 + .../_experimental/out/experimental/tag-management/index.html | 1 + litellm/proxy/_experimental/out/guardrails/index.html | 1 + litellm/proxy/_experimental/out/login/index.html | 1 + litellm/proxy/_experimental/out/logs/index.html | 1 + litellm/proxy/_experimental/out/mcp/oauth/callback/index.html | 1 + litellm/proxy/_experimental/out/model-hub/index.html | 1 + litellm/proxy/_experimental/out/model_hub/index.html | 1 + litellm/proxy/_experimental/out/model_hub_table/index.html | 1 + litellm/proxy/_experimental/out/models-and-endpoints/index.html | 1 + litellm/proxy/_experimental/out/onboarding/index.html | 1 + litellm/proxy/_experimental/out/organizations/index.html | 1 + litellm/proxy/_experimental/out/playground/index.html | 1 + litellm/proxy/_experimental/out/policies/index.html | 1 + .../proxy/_experimental/out/settings/admin-settings/index.html | 1 + .../_experimental/out/settings/logging-and-alerts/index.html | 1 + .../proxy/_experimental/out/settings/router-settings/index.html | 1 + litellm/proxy/_experimental/out/settings/ui-theme/index.html | 1 + litellm/proxy/_experimental/out/teams/index.html | 1 + litellm/proxy/_experimental/out/test-key/index.html | 1 + litellm/proxy/_experimental/out/tools/mcp-servers/index.html | 1 + litellm/proxy/_experimental/out/tools/vector-stores/index.html | 1 + litellm/proxy/_experimental/out/usage/index.html | 1 + litellm/proxy/_experimental/out/users/index.html | 1 + litellm/proxy/_experimental/out/virtual-keys/index.html | 1 + litellm/rag/ingestion/base_ingestion.py | 2 +- .../common_utils => litellm_core_utils}/test_url_utils.py | 2 +- 40 files changed, 39 insertions(+), 5 deletions(-) rename litellm/{proxy/common_utils => litellm_core_utils}/url_utils.py (100%) create mode 100644 litellm/proxy/_experimental/out/404/index.html create mode 100644 litellm/proxy/_experimental/out/_not-found/index.html create mode 100644 litellm/proxy/_experimental/out/api-reference/index.html create mode 100644 litellm/proxy/_experimental/out/chat/index.html create mode 100644 litellm/proxy/_experimental/out/experimental/api-playground/index.html create mode 100644 litellm/proxy/_experimental/out/experimental/budgets/index.html create mode 100644 litellm/proxy/_experimental/out/experimental/caching/index.html create mode 100644 litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html create mode 100644 litellm/proxy/_experimental/out/experimental/old-usage/index.html create mode 100644 litellm/proxy/_experimental/out/experimental/prompts/index.html create mode 100644 litellm/proxy/_experimental/out/experimental/tag-management/index.html create mode 100644 litellm/proxy/_experimental/out/guardrails/index.html create mode 100644 litellm/proxy/_experimental/out/login/index.html create mode 100644 litellm/proxy/_experimental/out/logs/index.html create mode 100644 litellm/proxy/_experimental/out/mcp/oauth/callback/index.html create mode 100644 litellm/proxy/_experimental/out/model-hub/index.html create mode 100644 litellm/proxy/_experimental/out/model_hub/index.html create mode 100644 litellm/proxy/_experimental/out/model_hub_table/index.html create mode 100644 litellm/proxy/_experimental/out/models-and-endpoints/index.html create mode 100644 litellm/proxy/_experimental/out/onboarding/index.html create mode 100644 litellm/proxy/_experimental/out/organizations/index.html create mode 100644 litellm/proxy/_experimental/out/playground/index.html create mode 100644 litellm/proxy/_experimental/out/policies/index.html create mode 100644 litellm/proxy/_experimental/out/settings/admin-settings/index.html create mode 100644 litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html create mode 100644 litellm/proxy/_experimental/out/settings/router-settings/index.html create mode 100644 litellm/proxy/_experimental/out/settings/ui-theme/index.html create mode 100644 litellm/proxy/_experimental/out/teams/index.html create mode 100644 litellm/proxy/_experimental/out/test-key/index.html create mode 100644 litellm/proxy/_experimental/out/tools/mcp-servers/index.html create mode 100644 litellm/proxy/_experimental/out/tools/vector-stores/index.html create mode 100644 litellm/proxy/_experimental/out/usage/index.html create mode 100644 litellm/proxy/_experimental/out/users/index.html create mode 100644 litellm/proxy/_experimental/out/virtual-keys/index.html rename tests/test_litellm/{proxy/common_utils => litellm_core_utils}/test_url_utils.py (97%) diff --git a/litellm/litellm_core_utils/prompt_templates/image_handling.py b/litellm/litellm_core_utils/prompt_templates/image_handling.py index c036ec23ddb..fd38bc9388d 100644 --- a/litellm/litellm_core_utils/prompt_templates/image_handling.py +++ b/litellm/litellm_core_utils/prompt_templates/image_handling.py @@ -10,7 +10,7 @@ import litellm from litellm import verbose_logger from litellm.caching.caching import InMemoryCache from litellm.constants import MAX_IMAGE_URL_DOWNLOAD_SIZE_MB -from litellm.proxy.common_utils.url_utils import async_safe_get, safe_get +from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get MAX_IMGS_IN_MEMORY = 10 diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 6245828a381..01e5dc39a34 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -30,7 +30,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.default_encoding import encoding as default_encoding from litellm.llms.custom_httpx.http_handler import _get_httpx_client -from litellm.proxy.common_utils.url_utils import safe_get +from litellm.litellm_core_utils.url_utils import safe_get from litellm.types.llms.anthropic import ( AnthropicMessagesToolResultParam, AnthropicMessagesToolUseParam, diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py similarity index 100% rename from litellm/proxy/common_utils/url_utils.py rename to litellm/litellm_core_utils/url_utils.py diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index d6b2cb86b26..3b2fa097b70 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -15,7 +15,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy.common_utils.url_utils import async_safe_get +from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) diff --git a/litellm/proxy/_experimental/out/404/index.html b/litellm/proxy/_experimental/out/404/index.html new file mode 100644 index 00000000000..344481d3aed --- /dev/null +++ b/litellm/proxy/_experimental/out/404/index.html @@ -0,0 +1 @@ +404: This page could not be found.LiteLLM Dashboard

404

This page could not be found.

\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/_not-found/index.html b/litellm/proxy/_experimental/out/_not-found/index.html new file mode 100644 index 00000000000..344481d3aed --- /dev/null +++ b/litellm/proxy/_experimental/out/_not-found/index.html @@ -0,0 +1 @@ +404: This page could not be found.LiteLLM Dashboard

404

This page could not be found.

\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/api-reference/index.html b/litellm/proxy/_experimental/out/api-reference/index.html new file mode 100644 index 00000000000..b636faba290 --- /dev/null +++ b/litellm/proxy/_experimental/out/api-reference/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/chat/index.html b/litellm/proxy/_experimental/out/chat/index.html new file mode 100644 index 00000000000..0d684c66cb5 --- /dev/null +++ b/litellm/proxy/_experimental/out/chat/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/api-playground/index.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html new file mode 100644 index 00000000000..5268cc3d9ca --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/api-playground/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/budgets/index.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html new file mode 100644 index 00000000000..f463b2d5df3 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/budgets/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/caching/index.html b/litellm/proxy/_experimental/out/experimental/caching/index.html new file mode 100644 index 00000000000..cf2a1aa14a0 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/caching/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html new file mode 100644 index 00000000000..069f97b082a --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/old-usage/index.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html new file mode 100644 index 00000000000..53540d126c4 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/old-usage/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/prompts/index.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html new file mode 100644 index 00000000000..615f06b8166 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/prompts/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/tag-management/index.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html new file mode 100644 index 00000000000..e7d0631c339 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/tag-management/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/guardrails/index.html b/litellm/proxy/_experimental/out/guardrails/index.html new file mode 100644 index 00000000000..ebbe174662b --- /dev/null +++ b/litellm/proxy/_experimental/out/guardrails/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/login/index.html b/litellm/proxy/_experimental/out/login/index.html new file mode 100644 index 00000000000..54472c6cc11 --- /dev/null +++ b/litellm/proxy/_experimental/out/login/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
🚅 LiteLLM
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/logs/index.html b/litellm/proxy/_experimental/out/logs/index.html new file mode 100644 index 00000000000..ec43b677a2f --- /dev/null +++ b/litellm/proxy/_experimental/out/logs/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html new file mode 100644 index 00000000000..830060c7aa2 --- /dev/null +++ b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/model-hub/index.html b/litellm/proxy/_experimental/out/model-hub/index.html new file mode 100644 index 00000000000..506c3695285 --- /dev/null +++ b/litellm/proxy/_experimental/out/model-hub/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/model_hub/index.html b/litellm/proxy/_experimental/out/model_hub/index.html new file mode 100644 index 00000000000..27bac5cde7d --- /dev/null +++ b/litellm/proxy/_experimental/out/model_hub/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/model_hub_table/index.html b/litellm/proxy/_experimental/out/model_hub_table/index.html new file mode 100644 index 00000000000..db5d0e6a718 --- /dev/null +++ b/litellm/proxy/_experimental/out/model_hub_table/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/models-and-endpoints/index.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html new file mode 100644 index 00000000000..96c1a43a7c0 --- /dev/null +++ b/litellm/proxy/_experimental/out/models-and-endpoints/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/onboarding/index.html b/litellm/proxy/_experimental/out/onboarding/index.html new file mode 100644 index 00000000000..5c2121443f9 --- /dev/null +++ b/litellm/proxy/_experimental/out/onboarding/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/organizations/index.html b/litellm/proxy/_experimental/out/organizations/index.html new file mode 100644 index 00000000000..51dd7d1c764 --- /dev/null +++ b/litellm/proxy/_experimental/out/organizations/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/playground/index.html b/litellm/proxy/_experimental/out/playground/index.html new file mode 100644 index 00000000000..41ef863e95b --- /dev/null +++ b/litellm/proxy/_experimental/out/playground/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/policies/index.html b/litellm/proxy/_experimental/out/policies/index.html new file mode 100644 index 00000000000..a452ae4c4aa --- /dev/null +++ b/litellm/proxy/_experimental/out/policies/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/admin-settings/index.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html new file mode 100644 index 00000000000..b29b4856b0d --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/admin-settings/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html new file mode 100644 index 00000000000..7d5d218fda4 --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/router-settings/index.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html new file mode 100644 index 00000000000..eb3fd3fde00 --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/router-settings/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/ui-theme/index.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html new file mode 100644 index 00000000000..17d352321c5 --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/ui-theme/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/teams/index.html b/litellm/proxy/_experimental/out/teams/index.html new file mode 100644 index 00000000000..781441c0732 --- /dev/null +++ b/litellm/proxy/_experimental/out/teams/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/test-key/index.html b/litellm/proxy/_experimental/out/test-key/index.html new file mode 100644 index 00000000000..22c06d24381 --- /dev/null +++ b/litellm/proxy/_experimental/out/test-key/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers/index.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html new file mode 100644 index 00000000000..64b747528e0 --- /dev/null +++ b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/tools/vector-stores/index.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html new file mode 100644 index 00000000000..098b0d212c6 --- /dev/null +++ b/litellm/proxy/_experimental/out/tools/vector-stores/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/usage/index.html b/litellm/proxy/_experimental/out/usage/index.html new file mode 100644 index 00000000000..ed6ac2eba97 --- /dev/null +++ b/litellm/proxy/_experimental/out/usage/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/users/index.html b/litellm/proxy/_experimental/out/users/index.html new file mode 100644 index 00000000000..247dda941bd --- /dev/null +++ b/litellm/proxy/_experimental/out/users/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/virtual-keys/index.html b/litellm/proxy/_experimental/out/virtual-keys/index.html new file mode 100644 index 00000000000..b17ef6de095 --- /dev/null +++ b/litellm/proxy/_experimental/out/virtual-keys/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index 6f139764a85..6a4eb89d0fd 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy.common_utils.url_utils import async_safe_get +from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.rag.ingestion.file_parsers import extract_text_from_pdf from litellm.rag.text_splitters import RecursiveCharacterTextSplitter from litellm.types.rag import RAGIngestOptions, RAGIngestResponse diff --git a/tests/test_litellm/proxy/common_utils/test_url_utils.py b/tests/test_litellm/litellm_core_utils/test_url_utils.py similarity index 97% rename from tests/test_litellm/proxy/common_utils/test_url_utils.py rename to tests/test_litellm/litellm_core_utils/test_url_utils.py index 4dbda6a815e..16798cebad5 100644 --- a/tests/test_litellm/proxy/common_utils/test_url_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_url_utils.py @@ -1,7 +1,7 @@ import pytest import litellm -from litellm.proxy.common_utils.url_utils import SSRFError, _is_blocked_ip, validate_url +from litellm.litellm_core_utils.url_utils import SSRFError, _is_blocked_ip, validate_url class TestIsBlockedIp: From 30c6556782103ece93b86a50b5ae64ffa7c887d6 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 05:06:14 +0000 Subject: [PATCH 16/41] test: bypass SSRF validation in image handling tests --- .../litellm_core_utils/test_image_handling.py | 39 +++++++++++++------ 1 file changed, 27 insertions(+), 12 deletions(-) diff --git a/tests/test_litellm/litellm_core_utils/test_image_handling.py b/tests/test_litellm/litellm_core_utils/test_image_handling.py index 9c2939b2da5..cc13e816dde 100644 --- a/tests/test_litellm/litellm_core_utils/test_image_handling.py +++ b/tests/test_litellm/litellm_core_utils/test_image_handling.py @@ -5,11 +5,22 @@ from httpx import Request, Response import litellm from litellm import constants +from litellm.litellm_core_utils.prompt_templates import image_handling from litellm.litellm_core_utils.prompt_templates.image_handling import ( convert_url_to_base64, ) +@pytest.fixture(autouse=True) +def _bypass_ssrf(monkeypatch): + """Bypass SSRF validation in image handling tests — tests use fake URLs.""" + monkeypatch.setattr( + image_handling, + "safe_get", + lambda client, url, **kw: client.get(url, follow_redirects=True), + ) + + class DummyClient: def get(self, url, follow_redirects=True): return Response(status_code=404, request=Request("GET", url)) @@ -37,9 +48,7 @@ def test_completion_with_invalid_image_url(monkeypatch): } ] with pytest.raises(litellm.ImageFetchError) as excinfo: - litellm.completion( - model="gemini/gemini-pro", messages=messages, api_key="test" - ) + litellm.completion(model="gemini/gemini-pro", messages=messages, api_key="test") assert excinfo.value.status_code == 400 assert "Unable to fetch image" in str(excinfo.value) @@ -81,7 +90,7 @@ class StreamingLargeImageClient: headers = {"Content-Type": "image/jpeg"} if self.include_content_length: headers["Content-Length"] = str(size_bytes) - + # Create a generator that yields chunks without creating the whole file in memory def generate_chunks(total_size, chunk_size=8192): bytes_sent = 0 @@ -89,7 +98,7 @@ class StreamingLargeImageClient: chunk = b"x" * min(chunk_size, total_size - bytes_sent) bytes_sent += len(chunk) yield chunk - + # Create response with streaming content response = Response( status_code=200, @@ -97,7 +106,9 @@ class StreamingLargeImageClient: request=Request("GET", url), ) # Mock the iter_bytes method to return our generator - response.iter_bytes = lambda chunk_size=8192: generate_chunks(size_bytes, chunk_size) + response.iter_bytes = lambda chunk_size=8192: generate_chunks( + size_bytes, chunk_size + ) return response @@ -121,7 +132,9 @@ def test_image_exceeds_size_limit_without_content_length(monkeypatch): This uses the old non-streaming mock for backward compatibility. """ monkeypatch.setattr( - litellm, "module_level_client", LargeImageClient(size_mb=100, include_content_length=False) + litellm, + "module_level_client", + LargeImageClient(size_mb=100, include_content_length=False), ) with pytest.raises(litellm.ImageFetchError) as excinfo: @@ -134,7 +147,7 @@ def test_streaming_download_protects_against_huge_files(monkeypatch): """ Test that streaming download aborts early when file exceeds size limit, preventing memory exhaustion from huge files (e.g., petabyte-sized files). - + This test verifies that the streaming implementation doesn't download the entire file into memory before checking size. Instead, it should abort as soon as the limit is exceeded during streaming. @@ -148,7 +161,7 @@ def test_streaming_download_protects_against_huge_files(monkeypatch): # Verify the error message shows it was caught during streaming assert "exceeds maximum allowed size" in str(excinfo.value) - + # The error should be raised after downloading just slightly more than the limit # not after downloading the full 1GB @@ -187,13 +200,15 @@ def test_streaming_download_handles_petabyte_file(monkeypatch): """ Test that streaming download can handle extremely large file URLs (e.g., petabyte-sized) without attempting to download the entire file or causing memory exhaustion. - + This simulates what happens if a malicious actor or misconfiguration provides a URL to an extremely large file. """ # Simulate a 1 petabyte file (1,000,000 GB) # Without streaming protection, this would cause OOM or hang indefinitely - client = StreamingLargeImageClient(size_mb=1_000_000_000, include_content_length=False) + client = StreamingLargeImageClient( + size_mb=1_000_000_000, include_content_length=False + ) monkeypatch.setattr(litellm, "module_level_client", client) with pytest.raises(litellm.ImageFetchError) as excinfo: @@ -214,6 +229,6 @@ def test_image_size_limit_disabled(monkeypatch): with pytest.raises(litellm.ImageFetchError) as excinfo: convert_url_to_base64("https://example.com/image.jpg") - + assert "Image URL download is disabled" in str(excinfo.value) assert "MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0" in str(excinfo.value) From f5a9218cb31bd1284a92bcfb70228ad0e51cbd97 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 05:14:39 +0000 Subject: [PATCH 17/41] chore: remove unused asyncio import --- litellm/litellm_core_utils/url_utils.py | 1 - 1 file changed, 1 deletion(-) diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index a47d7ae2ca4..1a552e6cf54 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -9,7 +9,6 @@ URL to connect to the validated IP directly — no TOCTOU gap, no DNS rebinding. Redirects are followed manually with validation at each hop. """ -import asyncio import socket from ipaddress import ip_address, ip_network from typing import Any, Tuple From 1f50c6fa66b0cb6cbe970c442e3ba4130cb8d642 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 21:28:13 +0000 Subject: [PATCH 18/41] test: mock DNS resolution, hoist httpx import to module level Greptile P1: six tests in test_url_utils.py performed real DNS lookups to example.com, violating the tests/test_litellm/ mock-only rule and risking offline CI failures. Add mock_dns_public and mock_dns_failure fixtures that monkeypatch socket.getaddrinfo on the url_utils module. Greptile P2: move 'import httpx' from inside _extract_redirect_url to module-level imports per CLAUDE.md style guide. --- litellm/litellm_core_utils/url_utils.py | 4 +- .../litellm_core_utils/test_url_utils.py | 41 ++++++++++++++++--- 2 files changed, 37 insertions(+), 8 deletions(-) diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index 1a552e6cf54..aaeb2bee7ef 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -14,6 +14,8 @@ from ipaddress import ip_address, ip_network from typing import Any, Tuple from urllib.parse import urlparse, urlunparse +import httpx + import litellm _BLOCKED_NETWORKS = [ @@ -142,8 +144,6 @@ _MAX_REDIRECTS = 10 def _extract_redirect_url(response: Any, request_url: str) -> str: """Extract and resolve the redirect target from a response's Location header.""" - import httpx - location = response.headers.get("location") if not location: raise SSRFError("Redirect response has no Location header") diff --git a/tests/test_litellm/litellm_core_utils/test_url_utils.py b/tests/test_litellm/litellm_core_utils/test_url_utils.py index 16798cebad5..1b8121efaca 100644 --- a/tests/test_litellm/litellm_core_utils/test_url_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_url_utils.py @@ -1,9 +1,34 @@ +import socket + import pytest import litellm +from litellm.litellm_core_utils import url_utils from litellm.litellm_core_utils.url_utils import SSRFError, _is_blocked_ip, validate_url +@pytest.fixture +def mock_dns_public(monkeypatch): + """Resolve any hostname to 93.184.216.34 (public).""" + + def fake_getaddrinfo(host, port, *args, **kwargs): + return [ + (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port or 80)) + ] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake_getaddrinfo) + + +@pytest.fixture +def mock_dns_failure(monkeypatch): + """Make every DNS lookup raise gaierror.""" + + def fake_getaddrinfo(host, port, *args, **kwargs): + raise socket.gaierror("Name or service not known") + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake_getaddrinfo) + + class TestIsBlockedIp: def test_blocks_private(self): assert _is_blocked_ip("10.0.0.1") is True @@ -48,22 +73,22 @@ class TestValidateUrl: with pytest.raises(SSRFError): validate_url("http:///path") - def test_allows_public_https(self): + def test_allows_public_https(self, mock_dns_public): rewritten, host = validate_url("https://example.com/image.png") assert host == "example.com" assert rewritten == "https://example.com/image.png" - def test_rewrites_public_http_to_ip(self): + def test_rewrites_public_http_to_ip(self, mock_dns_public): rewritten, host = validate_url("http://example.com/image.png") assert host == "example.com" assert "example.com" not in rewritten - def test_preserves_path_and_query(self): + def test_preserves_path_and_query(self, mock_dns_public): rewritten, host = validate_url("http://example.com/path?key=value") assert "/path" in rewritten assert "key=value" in rewritten - def test_dns_failure_raises(self): + def test_dns_failure_raises(self, mock_dns_failure): with pytest.raises(SSRFError, match="DNS resolution failed"): validate_url("http://this-domain-does-not-exist-xyz123.invalid/test") @@ -75,13 +100,17 @@ class TestValidateUrl: with pytest.raises(SSRFError): validate_url("http://[::1]/") - def test_https_rewrites_when_ssl_verify_disabled(self, monkeypatch): + def test_https_rewrites_when_ssl_verify_disabled( + self, monkeypatch, mock_dns_public + ): monkeypatch.setattr(litellm, "ssl_verify", False) rewritten, host = validate_url("https://example.com/image.png") assert host == "example.com" assert "example.com" not in rewritten # rewritten to IP - def test_https_not_rewritten_when_ssl_verify_enabled(self, monkeypatch): + def test_https_not_rewritten_when_ssl_verify_enabled( + self, monkeypatch, mock_dns_public + ): monkeypatch.setattr(litellm, "ssl_verify", True) rewritten, host = validate_url("https://example.com/image.png") assert rewritten == "https://example.com/image.png" From 22572eafaf3d59caec1da10b8cf9fc073abb5a29 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 21:29:13 +0000 Subject: [PATCH 19/41] fix: merge admin metadata from both metadata and litellm_metadata Greptile P2: _get_admin_metadata used 'litellm_metadata or metadata', meaning a caller sending a non-empty litellm_metadata would shadow admin config the proxy had injected into data['metadata']. Admin exemptions would be silently ignored. Check both keys and prefer whichever contains admin fields. Add regression test covering the shadowing scenario. --- litellm/integrations/custom_guardrail.py | 19 ++++++++++++++----- .../integrations/test_custom_guardrail.py | 14 ++++++++++++++ 2 files changed, 28 insertions(+), 5 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 89431847e90..b0931964cc6 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -257,11 +257,20 @@ class CustomGuardrail(CustomLogger): @staticmethod def _get_admin_metadata(data: dict) -> dict: - """Return merged admin-configured key and team metadata from the request data.""" - metadata = data.get("litellm_metadata") or data.get("metadata", {}) - team_meta = metadata.get("user_api_key_team_metadata") or {} - key_meta = metadata.get("user_api_key_metadata") or {} - # Key-level settings override team-level + """Return merged admin-configured key and team metadata from the request data. + + The proxy may inject admin metadata (user_api_key_metadata, + user_api_key_team_metadata) into either ``metadata`` or + ``litellm_metadata`` depending on endpoint. Check both so a caller + cannot shadow admin config by pre-populating the other key. + Key-level settings override team-level. + """ + team_meta: dict = {} + key_meta: dict = {} + for key in ("metadata", "litellm_metadata"): + meta = data.get(key) or {} + team_meta = meta.get("user_api_key_team_metadata") or team_meta + key_meta = meta.get("user_api_key_metadata") or key_meta return {**team_meta, **key_meta} def get_disable_global_guardrail(self, data: dict) -> Optional[bool]: diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 4c1ef853ab4..b8eb0857364 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -227,6 +227,20 @@ class TestCustomGuardrailShouldRunGuardrail: ) assert result is False, "Admin-configured disable should be respected" + # Test 5: Admin config in metadata isn't shadowed by user-supplied litellm_metadata + data_cross_key = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + "metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}}, + "litellm_metadata": {"request_tags": ["user-supplied"]}, + } + result = custom_guardrail.should_run_guardrail( + data=data_cross_key, event_type=GuardrailEventHooks.pre_call + ) + assert ( + result is False + ), "Admin config in metadata must not be shadowed by user-supplied litellm_metadata" + def test_should_run_guardrail_with_opted_out_global_guardrails(self): """Test that per-guardrail opt-out only works from admin metadata""" from litellm.types.guardrails import GuardrailEventHooks From 1d3dda93429c42e58c2fd9f0db0298cd743000e5 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 21:40:19 +0000 Subject: [PATCH 20/41] feat: add admin opt-out for user URL validation Two litellm-level flags wired through litellm_settings YAML: - user_url_validation (bool, default True): master switch. When False, safe_get/async_safe_get bypass validation and call client.get directly. - user_url_allowed_hosts (List[str], default []): per-host allowlist. Entries are 'host' (matches any port) or 'host:port' (port-specific). Matched hosts skip the blocked-networks check but still resolve DNS and still rewrite HTTP to the validated IP, preserving rebinding protection within the permitted name. Also fix an existing Host header bug: IPv6 literals (e.g. 2001:db8::1) were emitted unbracketed, producing ambiguous values like '2001:db8::1:8080' per RFC 7230 5.4. Bracket them consistently in _format_host_header. --- litellm/__init__.py | 2 + litellm/litellm_core_utils/url_utils.py | 79 ++++++-- .../litellm_core_utils/test_url_utils.py | 187 ++++++++++++++++++ 3 files changed, 253 insertions(+), 15 deletions(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index 3b67d9e0021..273af465b29 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -274,6 +274,8 @@ use_client: bool = False ssl_verify: Union[str, bool] = True ssl_security_level: Optional[str] = None ssl_certificate: Optional[str] = None +user_url_validation: bool = True +user_url_allowed_hosts: List[str] = [] ssl_ecdh_curve: Optional[ str ] = None # Set to 'X25519' to disable PQC and improve performance diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index aaeb2bee7ef..e920a044583 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -7,11 +7,21 @@ input (image_url, file_url, spec_path, etc.) to prevent SSRF attacks. validate_url() resolves DNS once, validates all IPs, and rewrites the URL to connect to the validated IP directly — no TOCTOU gap, no DNS rebinding. Redirects are followed manually with validation at each hop. + +Admins can opt out via two ``litellm`` globals (wired from proxy config): + +- ``litellm.user_url_validation`` (bool, default True): master switch. + When False, ``safe_get``/``async_safe_get`` perform a plain fetch with + no DNS check, no block list, and no rewrite. +- ``litellm.user_url_allowed_hosts`` (List[str], default []): per-host + allowlist. Entries are ``hostname`` or ``hostname:port`` (IPv6 hosts as + ``[addr]`` / ``[addr]:port``). Matching hosts skip the blocked-networks + check but still resolve DNS and still rewrite HTTP to the resolved IP. """ import socket from ipaddress import ip_address, ip_network -from typing import Any, Tuple +from typing import Any, List, Set, Tuple from urllib.parse import urlparse, urlunparse import httpx @@ -52,6 +62,36 @@ def _is_blocked_ip(addr: str) -> bool: return any(ip in net for net in _BLOCKED_NETWORKS) +def _normalize_host(host: str) -> str: + """Lowercase and strip a trailing dot from a hostname.""" + return host.lower().rstrip(".") + + +def _format_host_header(hostname: str, port: int, default_port: int) -> str: + """Build an RFC 7230 Host header value, bracketing IPv6 literals.""" + bracketed = f"[{hostname}]" if ":" in hostname else hostname + if port == default_port: + return bracketed + return f"{bracketed}:{port}" + + +def _is_host_allowlisted(hostname: str, effective_port: int) -> bool: + """Check whether a host is in the admin-configured allowlist. + + Admin entries may be ``hostname`` (any port) or ``hostname:port``. IPv6 + literals are written bracketed (``[::1]`` / ``[::1]:8080``). Matching + is case-insensitive on the hostname. + """ + configured: List[str] = getattr(litellm, "user_url_allowed_hosts", []) or [] + if not configured: + return False + normalized_host = _normalize_host(hostname) + host_repr = f"[{normalized_host}]" if ":" in normalized_host else normalized_host + candidates: Set[str] = {host_repr, f"{host_repr}:{effective_port}"} + allowlist: Set[str] = {_normalize_host(entry) for entry in configured if entry} + return bool(candidates & allowlist) + + def validate_url(url: str) -> Tuple[str, str]: """ Validate a user-supplied URL and rewrite it to connect to a validated IP. @@ -68,9 +108,9 @@ def validate_url(url: str) -> Tuple[str, str]: url: The user-supplied URL to validate. Returns: - Tuple of (rewritten_url, original_hostname). + Tuple of (rewritten_url, host_header). The rewritten URL has the hostname replaced with the validated IP. - The original hostname should be set as the Host header. + The host_header value should be sent as the Host header. Raises: SSRFError: If the URL scheme is invalid or the hostname resolves @@ -87,16 +127,15 @@ def validate_url(url: str) -> Tuple[str, str]: port = parsed.port default_port = 443 if parsed.scheme == "https" else 80 + effective_port = port if port is not None else default_port + host_header = _format_host_header(hostname, effective_port, default_port) - # Build the Host header value — include port when non-default - host_header = ( - hostname if (port is None or port == default_port) else f"{hostname}:{port}" - ) + is_allowlisted = _is_host_allowlisted(hostname, effective_port) # Resolve hostname and validate ALL addresses try: addrinfo = socket.getaddrinfo( - hostname, port or default_port, proto=socket.IPPROTO_TCP + hostname, effective_port, proto=socket.IPPROTO_TCP ) except socket.gaierror as e: raise SSRFError(f"DNS resolution failed for '{hostname}': {e}") @@ -104,13 +143,14 @@ def validate_url(url: str) -> Tuple[str, str]: if not addrinfo: raise SSRFError(f"No addresses found for '{hostname}'") - for family, type_, proto, canonname, sockaddr in addrinfo: - if _is_blocked_ip(sockaddr[0]): - raise SSRFError( - f"URL targets a blocked address ({sockaddr[0]}). " - "If this is a legitimate internal service, use a direct " - "provider configuration instead of a user-supplied URL." - ) + if not is_allowlisted: + for family, type_, proto, canonname, sockaddr in addrinfo: + if _is_blocked_ip(sockaddr[0]): + raise SSRFError( + f"URL targets a blocked address ({sockaddr[0]}). " + "If this is a legitimate internal service, add the host " + "to `user_url_allowed_hosts` in general_settings." + ) # For HTTPS with SSL verification enabled, TLS certificate validation # binds the connection to the hostname — DNS rebinding can't redirect @@ -159,6 +199,9 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: request. No DNS rebinding (resolve-and-rewrite). No redirect bypass (each hop validated). No breaking change for legitimate CDN redirects. + When ``litellm.user_url_validation`` is False, validation is bypassed + and this function delegates to ``client.get(url, follow_redirects=True)``. + Args: client: An httpx.Client (sync). url: The user-supplied URL. @@ -167,6 +210,9 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: Returns: The final httpx.Response. """ + if not getattr(litellm, "user_url_validation", True): + kwargs.setdefault("follow_redirects", True) + return client.get(url, **kwargs) kwargs.pop("follow_redirects", None) caller_headers = kwargs.pop("headers", {}) for _ in range(_MAX_REDIRECTS): @@ -185,6 +231,9 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any: """Async version of safe_get.""" + if not getattr(litellm, "user_url_validation", True): + kwargs.setdefault("follow_redirects", True) + return await client.get(url, **kwargs) kwargs.pop("follow_redirects", None) caller_headers = kwargs.pop("headers", {}) for _ in range(_MAX_REDIRECTS): diff --git a/tests/test_litellm/litellm_core_utils/test_url_utils.py b/tests/test_litellm/litellm_core_utils/test_url_utils.py index 1b8121efaca..f6282c6fe79 100644 --- a/tests/test_litellm/litellm_core_utils/test_url_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_url_utils.py @@ -114,3 +114,190 @@ class TestValidateUrl: monkeypatch.setattr(litellm, "ssl_verify", True) rewritten, host = validate_url("https://example.com/image.png") assert rewritten == "https://example.com/image.png" + + +class TestHostHeaderFormatting: + """RFC 7230 §5.4: IPv6 literals must be bracketed in the Host header.""" + + def test_ipv4_no_port(self, monkeypatch): + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + _, host = validate_url("http://example.com/") + assert host == "example.com" + + def test_ipv4_with_explicit_nondefault_port(self, monkeypatch): + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + _, host = validate_url("http://example.com:8080/") + assert host == "example.com:8080" + + def test_ipv4_with_explicit_default_port_strips_port(self, monkeypatch): + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + _, host = validate_url("http://example.com:80/") + assert host == "example.com" + + def test_ipv6_literal_is_bracketed_with_port(self, monkeypatch): + """Regression: IPv6 + port produced ambiguous `Host: 2001:db8::1:8080`.""" + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["[2001:db8::1]"]) + + def fake(host, port, *a, **kw): + return [ + ( + socket.AF_INET6, + socket.SOCK_STREAM, + 6, + "", + ("2001:db8::1", port, 0, 0), + ) + ] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + _, host = validate_url("http://[2001:db8::1]:8080/") + assert host == "[2001:db8::1]:8080" + + def test_ipv6_literal_is_bracketed_without_port(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["[2001:db8::1]"]) + + def fake(host, port, *a, **kw): + return [ + ( + socket.AF_INET6, + socket.SOCK_STREAM, + 6, + "", + ("2001:db8::1", port, 0, 0), + ) + ] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + _, host = validate_url("http://[2001:db8::1]/") + assert host == "[2001:db8::1]" + + +class TestValidationMasterSwitch: + def test_disabled_bypasses_fetch_in_safe_get(self, monkeypatch): + """When user_url_validation is False, safe_get delegates to client.get without validation.""" + monkeypatch.setattr(litellm, "user_url_validation", False) + + calls = [] + + class FakeClient: + def get(self, url, **kwargs): + calls.append((url, kwargs)) + + class R: + is_redirect = False + + return R() + + url_utils.safe_get(FakeClient(), "http://127.0.0.1/internal") + assert calls and calls[0][0] == "http://127.0.0.1/internal" + assert calls[0][1].get("follow_redirects") is True + + def test_enabled_still_blocks(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_validation", True) + with pytest.raises(SSRFError): + validate_url("http://127.0.0.1/") + + +class TestHostAllowlist: + def test_allowlisted_hostname_permits_private_ip(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + rewritten, host = validate_url("http://internal.corp/path") + assert host == "internal.corp" + assert "10.0.1.5" in rewritten + + def test_non_allowlisted_hostname_still_blocked(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + with pytest.raises(SSRFError): + validate_url("http://other.corp/") + + def test_allowlist_case_insensitive(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["Internal.Corp"]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + rewritten, _ = validate_url("http://internal.corp/") + assert "10.0.1.5" in rewritten + + def test_allowlist_with_port_matches_explicit_port(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:8080"]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + rewritten, host = validate_url("http://internal.corp:8080/") + assert host == "internal.corp:8080" + assert "10.0.1.5" in rewritten + + def test_allowlist_with_port_matches_default_port(self, monkeypatch): + """Admin entry `host:443` matches `https://host/` (port=None, default 443).""" + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:443"]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + # Should succeed — no SSRFError raised + validate_url("https://internal.corp/") + + def test_allowlist_port_specific_does_not_match_other_port(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:8080"]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + with pytest.raises(SSRFError): + validate_url("http://internal.corp:9090/") + + def test_allowlist_host_entry_matches_any_port(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + validate_url("http://internal.corp:9090/") + validate_url("https://internal.corp:8443/") + + def test_allowlist_permits_loopback(self, monkeypatch): + """Admin may opt into loopback if they explicitly configure it.""" + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["localhost"]) + # localhost resolves locally without needing mocks + rewritten, host = validate_url("http://localhost:8080/") + assert host == "localhost:8080" + + def test_empty_allowlist_retains_default_deny(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", []) + with pytest.raises(SSRFError): + validate_url("http://127.0.0.1/") + + def test_allowlist_strips_trailing_dot(self, monkeypatch): + monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp."]) + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + validate_url("http://internal.corp/") From d0601692b8dc5fff8f31714372f65c584a2abf20 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 21:48:36 +0000 Subject: [PATCH 21/41] fix(proxy): strip user_api_key_metadata injection slots from user input Expand the pre-call metadata strip to also remove user_api_key_metadata and user_api_key_team_metadata. The proxy writes these fields into data[_metadata_variable_name] with admin-authoritative values, but only into that one metadata key; the caller's value in the OTHER metadata key (metadata vs litellm_metadata) would otherwise persist and be picked up by _get_admin_metadata, letting a caller supply their own 'admin' config to disable guardrails, opt out of global policies, etc. VERIA-28 (High): Security Policy and Guardrail Bypass via Unsanitized Request Metadata. Add regression test at the proxy boundary verifying the strip, and extend the guardrail test to cover the post-strip admin-config path. --- litellm/proxy/litellm_pre_call_utils.py | 8 +- .../integrations/test_custom_guardrail.py | 16 ++ .../proxy/test_litellm_pre_call_utils.py | 225 +++++++++++++----- 3 files changed, 188 insertions(+), 61 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 95e1cbed44e..0b2c20b25a5 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -977,11 +977,17 @@ async def add_litellm_data_to_request( # noqa: PLR0915 "Setting client-provided x-api-key as api_key parameter (will override deployment key)" ) - # Strip internal pipeline state from user input + # Strip internal pipeline state and admin-injection slots from user input. + # The proxy writes user_api_key_metadata / user_api_key_team_metadata + # into data[_metadata_variable_name] below; if a caller pre-populates + # either key on the OTHER metadata field, _get_admin_metadata lookups + # would treat the caller's payload as admin-configured. for _meta_key in ("metadata", "litellm_metadata"): _user_meta = data.get(_meta_key) if isinstance(_user_meta, dict): _user_meta.pop("_pipeline_managed_guardrails", None) + _user_meta.pop("user_api_key_metadata", None) + _user_meta.pop("user_api_key_team_metadata", None) ########################################################## # Init - Proxy Server Request diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index b8eb0857364..0c904e9df50 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -241,6 +241,22 @@ class TestCustomGuardrailShouldRunGuardrail: result is False ), "Admin config in metadata must not be shadowed by user-supplied litellm_metadata" + # Test 6: After the pre-call strip runs, user-injected + # user_api_key_metadata in the non-authoritative metadata key is gone. + # _get_admin_metadata must then surface admin config unchanged. + data_post_strip = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + "metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}}, + "litellm_metadata": {}, # post-strip: attacker payload removed + } + result = custom_guardrail.should_run_guardrail( + data=data_post_strip, event_type=GuardrailEventHooks.pre_call + ) + assert ( + result is False + ), "Admin config in metadata must be respected when other metadata key is empty" + def test_should_run_guardrail_with_opted_out_global_guardrails(self): """Test that per-guardrail opt-out only works from admin metadata""" from litellm.types.guardrails import GuardrailEventHooks diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index cf7e71b14d4..e5adc673990 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -207,6 +207,79 @@ async def test_add_litellm_data_to_request_parses_string_metadata(): assert updated_data["metadata"]["generation_name"] == "gen123" +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_admin_injection_slots(): + """User-supplied user_api_key_metadata / user_api_key_team_metadata / + _pipeline_managed_guardrails must be stripped from both metadata keys + before the proxy writes its own admin-populated values. Otherwise a + caller can shadow admin config via the non-`_metadata_variable_name` + metadata key (e.g. litellm_metadata while the proxy writes to metadata). + """ + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + # Caller tries to inject admin config into BOTH metadata keys + attacker_admin_payload = {"disable_global_guardrails": True} + data = { + "model": "gpt-3.5-turbo", + "metadata": { + "user_api_key_metadata": attacker_admin_payload, + "user_api_key_team_metadata": attacker_admin_payload, + "_pipeline_managed_guardrails": ["evaded"], + }, + "litellm_metadata": { + "user_api_key_metadata": attacker_admin_payload, + "user_api_key_team_metadata": attacker_admin_payload, + "_pipeline_managed_guardrails": ["evaded"], + }, + } + + real_admin_metadata = {"admin_flag": "from_proxy"} + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata=real_admin_metadata, + team_metadata=real_admin_metadata, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + # The key that matches `_metadata_variable_name` gets proxy-populated + # with the real admin payload; the OTHER key must not retain the + # attacker's injection. + populated = updated["metadata"] + assert populated["user_api_key_metadata"] == real_admin_metadata + assert populated["user_api_key_team_metadata"] == real_admin_metadata + assert "_pipeline_managed_guardrails" not in populated or populated[ + "_pipeline_managed_guardrails" + ] != ["evaded"] + + other = updated.get("litellm_metadata") or {} + assert other.get("user_api_key_metadata") in (None, {}, real_admin_metadata) + assert other.get("user_api_key_team_metadata") in (None, {}, real_admin_metadata) + assert "_pipeline_managed_guardrails" not in other + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_user_spend_and_budget(): from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request @@ -221,7 +294,10 @@ async def test_add_litellm_data_to_request_user_spend_and_budget(): request_mock.client = MagicMock() request_mock.client.host = "127.0.0.1" - data = {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]} + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + } user_api_key_dict = UserAPIKeyAuth( api_key="hashed-key", @@ -1023,6 +1099,7 @@ def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data(): # Restore original model_group_settings litellm.model_group_settings = original_model_group_settings + import json import time from typing import Optional @@ -1040,15 +1117,16 @@ class TestCustomLogger(CustomLogger): def __init__(self): self.standard_logging_object: Optional[StandardLoggingPayload] = None super().__init__() - + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): print(f"SUCCESS CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}") self.standard_logging_object = kwargs.get("standard_logging_object") print(f"Captured standard_logging_object: {self.standard_logging_object}") - + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): print(f"FAILURE CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}") + @pytest.mark.asyncio async def test_add_litellm_metadata_from_request_headers(): """ @@ -1065,8 +1143,16 @@ async def test_add_litellm_metadata_from_request_headers(): try: # Prepare test data (ensure no streaming, add mock_response and api_key to route to litellm.acompletion) - headers = {"x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}'} - data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "stream": False, "mock_response": "Hi", "api_key": "fake-key"} + headers = { + "x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}' + } + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "stream": False, + "mock_response": "Hi", + "api_key": "fake-key", + } # Create mock request with headers mock_request = MagicMock(spec=Request) @@ -1078,9 +1164,7 @@ async def test_add_litellm_metadata_from_request_headers(): # Create mock user API key dict mock_user_api_key_dict = UserAPIKeyAuth( - api_key="test-key", - user_id="test-user", - org_id="test-org" + api_key="test-key", user_id="test-user", org_id="test-org" ) # Create mock proxy logging object @@ -1095,7 +1179,7 @@ async def test_add_litellm_metadata_from_request_headers(): async def mock_post_call_success_hook(*args, **kwargs): # Return the response unchanged - return kwargs.get('response', args[2] if len(args) > 2 else None) + return kwargs.get("response", args[2] if len(args) > 2 else None) mock_proxy_logging_obj.during_call_hook = mock_during_call_hook mock_proxy_logging_obj.pre_call_hook = mock_pre_call_hook @@ -1108,10 +1192,15 @@ async def test_add_litellm_metadata_from_request_headers(): general_settings = {} # Create mock select_data_generator with correct signature - def mock_select_data_generator(response=None, user_api_key_dict=None, request_data=None): + def mock_select_data_generator( + response=None, user_api_key_dict=None, request_data=None + ): async def mock_generator(): - yield "data: " + json.dumps({"choices": [{"delta": {"content": "Hello"}}]}) + "\n\n" + yield "data: " + json.dumps( + {"choices": [{"delta": {"content": "Hello"}}]} + ) + "\n\n" yield "data: [DONE]\n\n" + return mock_generator() # Create the processor @@ -1129,22 +1218,28 @@ async def test_add_litellm_metadata_from_request_headers(): select_data_generator=mock_select_data_generator, llm_router=None, model="gpt-4", - is_streaming_request=False + is_streaming_request=False, ) # Sleep for 3 seconds to allow logging to complete await asyncio.sleep(3) # Check if standard_logging_object was set - assert test_logger.standard_logging_object is not None, "standard_logging_object should be populated after LLM request" + assert ( + test_logger.standard_logging_object is not None + ), "standard_logging_object should be populated after LLM request" # Verify the logging object contains expected metadata standard_logging_obj = test_logger.standard_logging_object - print(f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}") + print( + f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}" + ) SPEND_LOGS_METADATA = standard_logging_obj["metadata"]["spend_logs_metadata"] - assert SPEND_LOGS_METADATA == dict(json.loads(headers["x-litellm-spend-logs-metadata"])), "spend_logs_metadata should be the same as the headers" + assert SPEND_LOGS_METADATA == dict( + json.loads(headers["x-litellm-spend-logs-metadata"]) + ), "spend_logs_metadata should be the same as the headers" finally: litellm.callbacks = original_callbacks @@ -1197,7 +1292,9 @@ def test_get_internal_user_header_from_mapping_returns_expected_header(): {"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"}, ] - header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings) + header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping( + mappings + ) assert header_name == "X-OpenWebUI-User-Id" @@ -1205,7 +1302,9 @@ def test_get_internal_user_header_from_mapping_none_when_absent(): mappings = [ {"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"} ] - header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings) + header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping( + mappings + ) assert header_name is None single = {"header_name": "X-Only-Customer", "litellm_user_role": "customer"} @@ -1218,7 +1317,10 @@ def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present(): headers = {"X-OpenWebUI-User-Id": "internal-user-123"} general_settings = { "user_header_mappings": [ - {"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}, + { + "header_name": "X-OpenWebUI-User-Id", + "litellm_user_role": "internal_user", + }, {"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"}, ] } @@ -1312,7 +1414,7 @@ async def test_team_guardrails_append_to_key_guardrails(): metadata = updated_data.get("metadata", {}) guardrails = metadata.get("guardrails", []) - + assert "key-guardrail-1" in guardrails assert "key-guardrail-2" in guardrails assert "team-guardrail-1" in guardrails @@ -1341,7 +1443,7 @@ async def test_request_guardrails_do_not_override_key_guardrails(): metadata={"guardrails": ["key-guardrail-1"]}, team_metadata={}, ) - + # Test case: Request with empty guardrails should not result in empty guardrails data_with_empty = { "model": "gpt-3.5-turbo", @@ -1361,7 +1463,7 @@ async def test_request_guardrails_do_not_override_key_guardrails(): _metadata = updated_data_empty.get("metadata", {}) requested_guardrails = _metadata.get("guardrails", []) - + assert "guardrails" not in updated_data_empty assert "key-guardrail-1" in requested_guardrails assert len(requested_guardrails) == 1 @@ -1476,7 +1578,10 @@ def test_update_model_if_key_alias_exists(): assert data["model"] == "xai/grok-4-fast-non-reasoning" # Test case 2: Key alias doesn't exist - data = {"model": "unknown-model", "messages": [{"role": "user", "content": "Hello"}]} + data = { + "model": "unknown-model", + "messages": [{"role": "user", "content": "Hello"}], + } user_api_key_dict = UserAPIKeyAuth( api_key="test-key", aliases={"modelAlias": "xai/grok-4-fast-non-reasoning"}, @@ -1594,16 +1699,22 @@ async def test_embedding_header_forwarding_with_model_group(): # Verify that only x- prefixed headers (except x-stainless) were forwarded forwarded_headers = updated_data["headers"] - assert "X-Custom-Header" in forwarded_headers, "X-Custom-Header should be forwarded" + assert ( + "X-Custom-Header" in forwarded_headers + ), "X-Custom-Header should be forwarded" assert forwarded_headers["X-Custom-Header"] == "custom-value" assert "X-Request-ID" in forwarded_headers, "X-Request-ID should be forwarded" assert forwarded_headers["X-Request-ID"] == "test-request-123" # Verify that authorization header was NOT forwarded (sensitive header) - assert "Authorization" not in forwarded_headers, "Authorization header should not be forwarded" + assert ( + "Authorization" not in forwarded_headers + ), "Authorization header should not be forwarded" # Verify that Content-Type was NOT forwarded (doesn't start with x-) - assert "Content-Type" not in forwarded_headers, "Content-Type should not be forwarded" + assert ( + "Content-Type" not in forwarded_headers + ), "Content-Type should not be forwarded" # Verify original data fields are preserved assert updated_data["model"] == "local-openai/text-embedding-3-small" @@ -1659,8 +1770,9 @@ async def test_embedding_header_forwarding_without_model_group_config(): ) # Verify that headers were NOT added since model is not in forward list - assert "headers" not in updated_data or updated_data.get("headers") is None, \ - "Headers should not be forwarded for models not in forward_client_headers_to_llm_api list" + assert ( + "headers" not in updated_data or updated_data.get("headers") is None + ), "Headers should not be forwarded for models not in forward_client_headers_to_llm_api list" # Verify original data fields are preserved assert updated_data["model"] == "text-embedding-ada-002" @@ -1714,7 +1826,9 @@ async def test_add_guardrails_from_policy_engine(): attachment_registry = get_attachment_registry() attachment_registry._attachments = [ PolicyAttachment(policy="global-baseline", scope="*"), # applies to all - PolicyAttachment(policy="healthcare", teams=["healthcare-team"]), # applies to healthcare team + PolicyAttachment( + policy="healthcare", teams=["healthcare-team"] + ), # applies to healthcare team ] attachment_registry._initialized = True @@ -1757,7 +1871,10 @@ async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_po data = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], - "policies": ["PII-POLICY-GLOBAL", "HIPAA-POLICY"], # Dynamic policies - should be accepted and removed + "policies": [ + "PII-POLICY-GLOBAL", + "HIPAA-POLICY", + ], # Dynamic policies - should be accepted and removed "metadata": {}, } @@ -1780,7 +1897,9 @@ async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_po ) # Verify that 'policies' was removed from the request body - assert "policies" not in data, "'policies' should be removed from request body to prevent forwarding to LLM provider" + assert ( + "policies" not in data + ), "'policies' should be removed from request body to prevent forwarding to LLM provider" # Verify that other fields are preserved assert "model" in data @@ -1869,7 +1988,9 @@ async def test_bearer_token_not_in_debug_logs(): from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request from litellm.proxy.proxy_server import ProxyConfig - secret_token = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.fakesignature" + secret_token = ( + "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.fakesignature" + ) mock_request = MagicMock(spec=Request) mock_request.headers = { @@ -1898,8 +2019,10 @@ async def test_bearer_token_not_in_debug_logs(): logger.setLevel(logging.DEBUG) try: - with patch("litellm.proxy.proxy_server.llm_router", None), \ - patch("litellm.proxy.proxy_server.premium_user", True): + with ( + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", True), + ): await add_litellm_data_to_request( data=data, request=mock_request, @@ -2020,9 +2143,7 @@ def test_resolve_project_model_specific_wins(): "gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}}, "defaultconfig": {"azure": {"litellm_credentials": "team-default"}}, } - result = _resolve_credential_from_model_config( - "gpt-4", project_config, team_config - ) + result = _resolve_credential_from_model_config("gpt-4", project_config, team_config) assert result == "proj-gpt4" @@ -2034,9 +2155,7 @@ def test_resolve_project_default_wins_over_team(): "gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}}, "defaultconfig": {"azure": {"litellm_credentials": "team-default"}}, } - result = _resolve_credential_from_model_config( - "gpt-4", project_config, team_config - ) + result = _resolve_credential_from_model_config("gpt-4", project_config, team_config) assert result == "proj-default" @@ -2091,12 +2210,8 @@ def test_apply_overrides_project_model_specific(setup_test_credentials): }, project_metadata={ "model_config": { - "defaultconfig": { - "azure": {"litellm_credentials": "hotel-rec-azure"} - }, - "gpt-4-vision": { - "azure": {"litellm_credentials": "hotel-rec-vision"} - }, + "defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}}, + "gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}}, } }, ) @@ -2123,12 +2238,8 @@ def test_apply_overrides_project_default(setup_test_credentials): }, project_metadata={ "model_config": { - "defaultconfig": { - "azure": {"litellm_credentials": "hotel-rec-azure"} - }, - "gpt-4-vision": { - "azure": {"litellm_credentials": "hotel-rec-vision"} - }, + "defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}}, + "gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}}, } }, ) @@ -2231,9 +2342,7 @@ def test_apply_overrides_missing_credential_name(setup_test_credentials): api_key="test-key", team_metadata={ "model_config": { - "gpt-4": { - "azure": {"litellm_credentials": "nonexistent-credential"} - } + "gpt-4": {"azure": {"litellm_credentials": "nonexistent-credential"}} } }, ) @@ -2272,9 +2381,7 @@ def test_apply_overrides_no_model_in_data(setup_test_credentials): api_key="test-key", team_metadata={ "model_config": { - "defaultconfig": { - "azure": {"litellm_credentials": "some-cred"} - } + "defaultconfig": {"azure": {"litellm_credentials": "some-cred"}} } }, ) @@ -2305,9 +2412,7 @@ def test_apply_overrides_clientside_api_version_preserved(setup_test_credentials api_key="test-key", team_metadata={ "model_config": { - "gpt-4-vision": { - "azure": {"litellm_credentials": "hotel-rec-vision"} - } + "gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}} } }, ) From 0602564b66bd20147f0e0d9fd3bc9e4c5e4fd221 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 22:03:47 +0000 Subject: [PATCH 22/41] fix: switch blocklist to RFC 6890 via ipaddress.is_global, block multicast and Azure Wire Server MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replace the hand-maintained _BLOCKED_NETWORKS CIDR list with a default-deny check based on ipaddress.is_global (RFC 6890 semantics, implemented by Python's stdlib). Also reject multicast explicitly — is_global returns True for public multicast allocations, which are not legitimate HTTP targets. Only globally-routable cloud-fabric IPs need explicit exceptions; the canonical list contains one entry today: Azure Wire Server (168.63.129.16), an in-fabric service reachable from any Azure VM. Coverage delta picked up automatically via is_global: - Alibaba Cloud metadata (100.100.100.200, CGNAT) - Legacy Oracle metadata (192.0.0.192, IETF Protocol Assignments) - IPv4 documentation ranges (192.0.2.0/24, 198.51.100.0/24, 203.0.113.0/24) - IPv4 reserved/future-use (240.0.0.0/4) and broadcast - IPv6 documentation (2001:db8::/32) Also fix two issues Greptile flagged: - HTTP relative-redirect hops lost the original hostname because _extract_redirect_url joined the Location against the rewritten (IP-based) URL. Join against the pre-rewrite URL so the next hop's Host header keeps the original hostname. - Two unit tests performed real socket.getaddrinfo('localhost') calls. Monkeypatch them. Add coverage tests for every cloud-metadata IP from the canonical SSRF dictionary (AWS/GCP/Azure/Alibaba/Oracle/DO/OpenStack) plus the new multicast/reserved/documentation/broadcast ranges, and a regression test for redirect-hostname preservation. --- litellm/litellm_core_utils/url_utils.py | 39 +++++--- .../litellm_core_utils/test_url_utils.py | 97 ++++++++++++++++++- 2 files changed, 118 insertions(+), 18 deletions(-) diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index e920a044583..d22f92642a3 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -28,19 +28,13 @@ import httpx import litellm -_BLOCKED_NETWORKS = [ - ip_network("0.0.0.0/8"), - ip_network("10.0.0.0/8"), - ip_network("100.64.0.0/10"), - ip_network("127.0.0.0/8"), - ip_network("169.254.0.0/16"), - ip_network("172.16.0.0/12"), - ip_network("192.0.0.0/24"), - ip_network("192.168.0.0/16"), - ip_network("198.18.0.0/15"), - ip_network("::1/128"), - ip_network("fc00::/7"), - ip_network("fe80::/10"), +# Globally-routable IPs that are cloud-internal. Everything else +# non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented by +# Python's ``ipaddress`` module). This list only holds IPs that are +# publicly routable *and* point to cloud-fabric services reachable from +# inside a VM via special in-fabric routing. +_CLOUD_METADATA_EXCEPTIONS = [ + ip_network("168.63.129.16/32"), # Azure Wire Server ] _ALLOWED_SCHEMES = ("http", "https") @@ -53,13 +47,22 @@ class SSRFError(ValueError): def _is_blocked_ip(addr: str) -> bool: + """Return True for any IP not safe to reach from a user-supplied URL. + + Policy: default-deny via ``ip.is_global`` (RFC 6890), plus an explicit + exception list for globally-routable cloud-fabric IPs that are still + dangerous from inside a cloud VM (currently just Azure Wire Server). + Unparseable addresses fail closed. + """ try: ip = ip_address(addr) except ValueError: return True # fail-closed: unparseable addresses are blocked if ip.version == 6 and hasattr(ip, "ipv4_mapped") and ip.ipv4_mapped: ip = ip.ipv4_mapped - return any(ip in net for net in _BLOCKED_NETWORKS) + if not ip.is_global or ip.is_multicast: + return True + return any(ip in net for net in _CLOUD_METADATA_EXCEPTIONS) def _normalize_host(host: str) -> str: @@ -225,7 +228,9 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any: ) if not response.is_redirect: return response - url = _extract_redirect_url(response, validated_url) + # Resolve the next hop against the ORIGINAL (pre-rewrite) URL so + # relative Location headers keep the original hostname. + url = _extract_redirect_url(response, url) raise SSRFError("Too many redirects") @@ -246,5 +251,7 @@ async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any: ) if not response.is_redirect: return response - url = _extract_redirect_url(response, validated_url) + # Resolve the next hop against the ORIGINAL (pre-rewrite) URL so + # relative Location headers keep the original hostname. + url = _extract_redirect_url(response, url) raise SSRFError("Too many redirects") diff --git a/tests/test_litellm/litellm_core_utils/test_url_utils.py b/tests/test_litellm/litellm_core_utils/test_url_utils.py index f6282c6fe79..4579c203218 100644 --- a/tests/test_litellm/litellm_core_utils/test_url_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_url_utils.py @@ -39,6 +39,46 @@ class TestIsBlockedIp: def test_unparseable_is_blocked(self): assert _is_blocked_ip("not-an-ip") is True + # Coverage delta picked up by switching to `not ip.is_global` (RFC 6890) + # over the old hand-maintained CIDR list. + def test_blocks_cgnat_alibaba_metadata(self): + """100.100.100.200 is Alibaba Cloud metadata; lives in CGNAT.""" + assert _is_blocked_ip("100.100.100.200") is True + + def test_blocks_ietf_protocol_assignments_old_oracle_metadata(self): + """192.0.0.192 was the legacy Oracle Cloud metadata IP.""" + assert _is_blocked_ip("192.0.0.192") is True + + def test_blocks_documentation_ranges(self): + assert _is_blocked_ip("192.0.2.1") is True + assert _is_blocked_ip("198.51.100.1") is True + assert _is_blocked_ip("203.0.113.1") is True + + def test_blocks_multicast(self): + assert _is_blocked_ip("224.0.0.1") is True + + def test_blocks_reserved_future_use(self): + assert _is_blocked_ip("240.0.0.1") is True + + def test_blocks_broadcast(self): + assert _is_blocked_ip("255.255.255.255") is True + + def test_blocks_azure_wire_server(self): + """168.63.129.16 is globally routable but cloud-internal — explicit exception.""" + assert _is_blocked_ip("168.63.129.16") is True + + def test_blocks_aws_ipv6_imds(self): + """fd00:ec2::254 is AWS's IPv6 IMDS, in IPv6 ULA (fc00::/7).""" + assert _is_blocked_ip("fd00:ec2::254") is True + + def test_blocks_ipv4_mapped_private(self): + """::ffff:10.0.0.1 must be unwrapped and blocked as 10.0.0.1.""" + assert _is_blocked_ip("::ffff:10.0.0.1") is True + + def test_blocks_ipv4_mapped_azure_wire_server(self): + """::ffff:168.63.129.16 must be unwrapped and blocked via the exception list.""" + assert _is_blocked_ip("::ffff:168.63.129.16") is True + class TestValidateUrl: def test_blocks_loopback(self): @@ -92,7 +132,13 @@ class TestValidateUrl: with pytest.raises(SSRFError, match="DNS resolution failed"): validate_url("http://this-domain-does-not-exist-xyz123.invalid/test") - def test_blocks_localhost_hostname(self): + def test_blocks_localhost_hostname(self, monkeypatch): + def fake(host, port, *a, **kw): + return [ + (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", port or 80)) + ] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) with pytest.raises(SSRFError): validate_url("http://localhost/") @@ -181,6 +227,49 @@ class TestHostHeaderFormatting: assert host == "[2001:db8::1]" +class TestRedirectHostnamePreservation: + """Relative-location redirects must keep the original hostname, not the + rewritten IP, so the next hop's Host header still identifies the site.""" + + def test_relative_redirect_preserves_hostname_for_next_hop(self, monkeypatch): + def fake(host, port, *a, **kw): + return [ + (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port)) + ] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) + + class FakeResponse: + def __init__(self, status, location=None): + self.status_code = status + self.headers = {"location": location} if location else {} + self.is_redirect = 300 <= status < 400 + + hops = [] + + class FakeClient: + def __init__(self): + self._n = 0 + + def get(self, url, headers=None, follow_redirects=False, **kw): + hops.append({"url": url, "host": (headers or {}).get("Host")}) + self._n += 1 + if self._n == 1: + return FakeResponse(302, "/redirected") + return FakeResponse(200) + + url_utils.safe_get(FakeClient(), "http://example.com/initial") + assert len(hops) == 2 + # Both hops must carry the ORIGINAL hostname in the Host header. + assert hops[0]["host"] == "example.com" + assert hops[1]["host"] == "example.com" + # Both outbound URLs go to the resolved IP (rewritten), not the hostname. + assert "93.184.216.34" in hops[0]["url"] + assert "93.184.216.34" in hops[1]["url"] + # The second hop resolved /redirected relative to the original, not the IP. + assert hops[1]["url"].endswith("/redirected") + + class TestValidationMasterSwitch: def test_disabled_bypasses_fetch_in_safe_get(self, monkeypatch): """When user_url_validation is False, safe_get delegates to client.get without validation.""" @@ -284,7 +373,11 @@ class TestHostAllowlist: def test_allowlist_permits_loopback(self, monkeypatch): """Admin may opt into loopback if they explicitly configure it.""" monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["localhost"]) - # localhost resolves locally without needing mocks + + def fake(host, port, *a, **kw): + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", port))] + + monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake) rewritten, host = validate_url("http://localhost:8080/") assert host == "localhost:8080" From aa2f05f8c9663202a0d9e68dc1980bb6134b1d82 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 22:25:24 +0000 Subject: [PATCH 23/41] style: use 'is not None' for port check (handle port 0 explicitly) --- litellm/litellm_core_utils/url_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index d22f92642a3..b55882819de 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -170,7 +170,7 @@ def validate_url(url: str) -> Tuple[str, str]: is_ipv6 = addrinfo[0][0] == socket.AF_INET6 ip_host = f"[{validated_ip}]" if is_ipv6 else validated_ip - if port: + if port is not None: new_netloc = f"{ip_host}:{port}" else: new_netloc = ip_host From 0e62addd947087a9b75307af39d23ee8173a11a2 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 22:31:00 +0000 Subject: [PATCH 24/41] fix(proxy): gate caller-supplied routing/budget tags behind allow_client_tags VERIA-28 (High) follow-up: tag-based routing and tag budget enforcement read metadata.tags directly from the request, letting an attacker reach restricted tag-routed deployments or misattribute spend to a victim team's tag. Strip metadata.tags (and litellm_metadata.tags) at the pre-call boundary unless the caller's key or team metadata opts in with allow_client_tags=True. Default-deny: existing clients that need to pass routing tags must have the flag set explicitly on their key or team. Preserves the tag-routing feature for admins who trust their callers; closes the injection path for everyone else. --- litellm/proxy/litellm_pre_call_utils.py | 22 +++ .../proxy/test_litellm_pre_call_utils.py | 135 ++++++++++++++++++ 2 files changed, 157 insertions(+) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 0b2c20b25a5..f7b66a55b8e 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -989,6 +989,28 @@ async def add_litellm_data_to_request( # noqa: PLR0915 _user_meta.pop("user_api_key_metadata", None) _user_meta.pop("user_api_key_team_metadata", None) + # Strip caller-supplied routing/budget tags unless the admin has opted + # this key or team in via metadata.allow_client_tags=True. Tags drive + # tag-based routing and tag budget attribution — accepting them from + # untrusted callers lets an attacker reach restricted deployments or + # misattribute spend to a victim team's tag. + _admin_allow_client_tags = False + for _admin_meta in ( + user_api_key_dict.metadata, + user_api_key_dict.team_metadata, + ): + if ( + isinstance(_admin_meta, dict) + and _admin_meta.get("allow_client_tags") is True + ): + _admin_allow_client_tags = True + break + if not _admin_allow_client_tags: + for _meta_key in ("metadata", "litellm_metadata"): + _user_meta = data.get(_meta_key) + if isinstance(_user_meta, dict): + _user_meta.pop("tags", None) + ########################################################## # Init - Proxy Server Request # we do this as soon as entering so we track the original request diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index e5adc673990..e433bebb3ac 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -280,6 +280,141 @@ async def test_add_litellm_data_to_request_strips_admin_injection_slots(): assert "_pipeline_managed_guardrails" not in other +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_user_tags_without_permission(): + """Caller-supplied metadata.tags must be stripped when the key/team + metadata does not opt in via allow_client_tags=True. Otherwise an + attacker can reach restricted tag-routed deployments or attribute + spend to a victim team's tag.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"tags": ["restricted-tier", "victim-team"]}, + "litellm_metadata": {"tags": ["also-stripped"]}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert "tags" not in (updated.get("metadata") or {}) + assert "tags" not in (updated.get("litellm_metadata") or {}) + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_preserves_user_tags_when_key_opts_in(): + """When key.metadata.allow_client_tags=True, caller-supplied tags are + preserved and reach the router.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"tags": ["opted-in-tag"]}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"allow_client_tags": True}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"].get("tags") == ["opted-in-tag"] + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_preserves_user_tags_when_team_opts_in(): + """Team-level allow_client_tags is also honored (not just key-level).""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"tags": ["team-allowed"]}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={"allow_client_tags": True}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"].get("tags") == ["team-allowed"] + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_user_spend_and_budget(): from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request From 8526628a8fac6400b75b5bdb78958a656fa3abd4 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 22:40:43 +0000 Subject: [PATCH 25/41] test: update tag-merge tests for default-deny client tag policy Two pre-existing tests codified the pre-fix behavior where any caller- supplied metadata.tags would flow through to spend logs and routing: - test_add_key_or_team_level_spend_logs_metadata_to_request exercised the request/key/team tag merge. Set allow_client_tags=True on the key metadata so the merge path is still tested under the new regime. - test_create_file_with_nested_litellm_metadata asserted that litellm_metadata[tags] form-data propagated to the handler. Drop the tag field; the test still proves nested form-parser correctness via spend_logs_metadata and environment. --- tests/proxy_unit_tests/test_proxy_utils.py | 4 ++++ .../proxy/openai_files_endpoint/test_files_endpoint.py | 6 +++--- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 9f5f14457e8..a84bf7a4e7c 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -162,9 +162,13 @@ async def test_add_key_or_team_level_spend_logs_metadata_to_request( print(f"team_sl_metadata: {team_sl_metadata}") mock_request.url.path = "/chat/completions" + # Opt the key into client-supplied tags so request_tags are preserved + # and merged with admin-configured key/team tags. Without this flag, + # request_tags would be stripped by add_litellm_data_to_request. key_metadata = { "tags": key_tags, "spend_logs_metadata": key_sl_metadata, + "allow_client_tags": True, } team_metadata = { "tags": team_tags, diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 09d11388d84..49fff6de0ff 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -1295,7 +1295,6 @@ def test_create_file_with_nested_litellm_metadata( "target_model_names": "gpt-3.5-turbo", "litellm_metadata[spend_logs_metadata][owner]": "john_doe", "litellm_metadata[spend_logs_metadata][team]": "engineering", - "litellm_metadata[tags]": "production", "litellm_metadata[environment]": "prod", }, headers={"Authorization": "Bearer test-key"}, @@ -1306,11 +1305,12 @@ def test_create_file_with_nested_litellm_metadata( result = response.json() assert result["id"] == "file-test-123" - # Verify nested metadata was correctly parsed + # Verify nested metadata was correctly parsed. + # Note: caller-supplied `tags` is stripped by default; test removed + # to keep the parsing test focused on parser correctness. assert "spend_logs_metadata" in captured_litellm_metadata assert captured_litellm_metadata["spend_logs_metadata"]["owner"] == "john_doe" assert captured_litellm_metadata["spend_logs_metadata"]["team"] == "engineering" - assert captured_litellm_metadata["tags"] == "production" assert captured_litellm_metadata["environment"] == "prod" From af8d479482d8c41e8abc17e144c68b92e0d82d64 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 22:42:14 +0000 Subject: [PATCH 26/41] chore(proxy): emit warning when caller-supplied tags are stripped Silent strip is the worst debug UX: admin's client sends routing tags, they disappear, admin can't figure out why. Emit a warning naming the metadata key the tags came from and telling the admin exactly which flag to set if this is intentional. --- litellm/proxy/litellm_pre_call_utils.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index f7b66a55b8e..b1dae22c9fe 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1006,10 +1006,19 @@ async def add_litellm_data_to_request( # noqa: PLR0915 _admin_allow_client_tags = True break if not _admin_allow_client_tags: + _stripped_from: List[str] = [] for _meta_key in ("metadata", "litellm_metadata"): _user_meta = data.get(_meta_key) - if isinstance(_user_meta, dict): + if isinstance(_user_meta, dict) and "tags" in _user_meta: _user_meta.pop("tags", None) + _stripped_from.append(_meta_key) + if _stripped_from: + verbose_proxy_logger.warning( + "Stripped caller-supplied tags from %s: this key/team does " + "not have `allow_client_tags: true` in its metadata. Set it " + "to opt into client-supplied routing/budget tags.", + ", ".join(_stripped_from), + ) ########################################################## # Init - Proxy Server Request From 9622864ef20048065c19cd0992e5acb33c97c67e Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 22:50:54 +0000 Subject: [PATCH 27/41] test: set allow_client_tags on duplicate-tag merge test MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit test_add_litellm_data_to_request_duplicate_tags tests the request/key tag merge when tags overlap. The merge requires caller-supplied tags to flow through — set allow_client_tags=True on the key so the merge path stays testable under the new default-deny regime. --- tests/proxy_unit_tests/test_proxy_utils.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index a84bf7a4e7c..916e01b0f2c 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -842,12 +842,13 @@ async def test_add_litellm_data_to_request_duplicate_tags( mock_request.headers = {} mock_request.state = State() - # Setup key with tags in metadata + # Setup key with tags in metadata. Opt into client-supplied tags so the + # request_tags are preserved for the merge under test. user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", user_id="test_user_id", org_id="test_org_id", - metadata={"tags": key_tags}, + metadata={"tags": key_tags, "allow_client_tags": True}, ) # Setup request data with tags From db19f24d694ec82037f06361e4aa6ae10d194da5 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 23:00:52 +0000 Subject: [PATCH 28/41] fix(proxy): move metadata strip after JSON-string parse MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Veria AI caught a bypass: metadata can arrive as a JSON string via multipart/form-data or extra_body, and the existing strip block ran before the string-to-dict parse. The isinstance(_user_meta, dict) guard returned False on the string, the strip was skipped, and then the parse turned the string into a dict — leaving user_api_key_metadata / user_api_key_team_metadata / _pipeline_managed_guardrails / tags intact in the parsed dict. Move the strip to run AFTER the parse and BEFORE the merge of litellm_metadata into data[_metadata_variable_name], closing the bypass for both raw-dict and string-encoded payloads. Regression test: test_add_litellm_data_to_request_strips_string_encoded_admin_injection. --- litellm/proxy/litellm_pre_call_utils.py | 103 ++++++++++-------- .../proxy/test_litellm_pre_call_utils.py | 68 +++++++++++- 2 files changed, 122 insertions(+), 49 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index b1dae22c9fe..33affa5351c 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -977,49 +977,6 @@ async def add_litellm_data_to_request( # noqa: PLR0915 "Setting client-provided x-api-key as api_key parameter (will override deployment key)" ) - # Strip internal pipeline state and admin-injection slots from user input. - # The proxy writes user_api_key_metadata / user_api_key_team_metadata - # into data[_metadata_variable_name] below; if a caller pre-populates - # either key on the OTHER metadata field, _get_admin_metadata lookups - # would treat the caller's payload as admin-configured. - for _meta_key in ("metadata", "litellm_metadata"): - _user_meta = data.get(_meta_key) - if isinstance(_user_meta, dict): - _user_meta.pop("_pipeline_managed_guardrails", None) - _user_meta.pop("user_api_key_metadata", None) - _user_meta.pop("user_api_key_team_metadata", None) - - # Strip caller-supplied routing/budget tags unless the admin has opted - # this key or team in via metadata.allow_client_tags=True. Tags drive - # tag-based routing and tag budget attribution — accepting them from - # untrusted callers lets an attacker reach restricted deployments or - # misattribute spend to a victim team's tag. - _admin_allow_client_tags = False - for _admin_meta in ( - user_api_key_dict.metadata, - user_api_key_dict.team_metadata, - ): - if ( - isinstance(_admin_meta, dict) - and _admin_meta.get("allow_client_tags") is True - ): - _admin_allow_client_tags = True - break - if not _admin_allow_client_tags: - _stripped_from: List[str] = [] - for _meta_key in ("metadata", "litellm_metadata"): - _user_meta = data.get(_meta_key) - if isinstance(_user_meta, dict) and "tags" in _user_meta: - _user_meta.pop("tags", None) - _stripped_from.append(_meta_key) - if _stripped_from: - verbose_proxy_logger.warning( - "Stripped caller-supplied tags from %s: this key/team does " - "not have `allow_client_tags: true` in its metadata. Set it " - "to opt into client-supplied routing/budget tags.", - ", ".join(_stripped_from), - ) - ########################################################## # Init - Proxy Server Request # we do this as soon as entering so we track the original request @@ -1126,11 +1083,61 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ) else: data["litellm_metadata"] = parsed_litellm_metadata - # Merge litellm_metadata into the metadata variable (preserving existing values) - if isinstance(data["litellm_metadata"], dict): - for key, value in data["litellm_metadata"].items(): - if key not in data[_metadata_variable_name]: - data[_metadata_variable_name][key] = value + + # Strip internal pipeline state and admin-injection slots from user input. + # Runs AFTER the string-to-dict parse above so JSON-string metadata (sent + # via multipart/form-data or extra_body) cannot smuggle `user_api_key_metadata` + # past the isinstance(dict) guard. + # + # The proxy writes user_api_key_metadata / user_api_key_team_metadata into + # data[_metadata_variable_name] below; if a caller pre-populates either + # key on the OTHER metadata field, _get_admin_metadata lookups would treat + # the caller's payload as admin-configured. + for _meta_key in ("metadata", "litellm_metadata"): + _user_meta = data.get(_meta_key) + if isinstance(_user_meta, dict): + _user_meta.pop("_pipeline_managed_guardrails", None) + _user_meta.pop("user_api_key_metadata", None) + _user_meta.pop("user_api_key_team_metadata", None) + + # Strip caller-supplied routing/budget tags unless the admin has opted + # this key or team in via metadata.allow_client_tags=True. Tags drive + # tag-based routing and tag budget attribution — accepting them from + # untrusted callers lets an attacker reach restricted deployments or + # misattribute spend to a victim team's tag. + _admin_allow_client_tags = False + for _admin_meta in ( + user_api_key_dict.metadata, + user_api_key_dict.team_metadata, + ): + if ( + isinstance(_admin_meta, dict) + and _admin_meta.get("allow_client_tags") is True + ): + _admin_allow_client_tags = True + break + if not _admin_allow_client_tags: + _stripped_from: List[str] = [] + for _meta_key in ("metadata", "litellm_metadata"): + _user_meta = data.get(_meta_key) + if isinstance(_user_meta, dict) and "tags" in _user_meta: + _user_meta.pop("tags", None) + _stripped_from.append(_meta_key) + if _stripped_from: + verbose_proxy_logger.warning( + "Stripped caller-supplied tags from %s: this key/team does " + "not have `allow_client_tags: true` in its metadata. Set it " + "to opt into client-supplied routing/budget tags.", + ", ".join(_stripped_from), + ) + + # Now merge litellm_metadata into the metadata variable (preserving existing + # values) — runs AFTER the strip so attacker injections in litellm_metadata + # cannot cross-contaminate the admin-authoritative metadata dict. + if "litellm_metadata" in data and isinstance(data["litellm_metadata"], dict): + for key, value in data["litellm_metadata"].items(): + if key not in data[_metadata_variable_name]: + data[_metadata_variable_name][key] = value data = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( data=data, diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index e433bebb3ac..351857b9b10 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -280,6 +280,70 @@ async def test_add_litellm_data_to_request_strips_admin_injection_slots(): assert "_pipeline_managed_guardrails" not in other +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection(): + """Regression: metadata arriving as a JSON string (multipart/form-data or + extra_body) must not bypass the admin-injection strip. The parse happens + AFTER receipt, so the strip has to run after the parse, not before. + """ + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "multipart/form-data"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + # Attacker encodes an admin-injection payload inside a JSON string. + attacker_payload = { + "user_api_key_metadata": {"disable_global_guardrails": True}, + "user_api_key_team_metadata": {"disable_global_guardrails": True}, + "_pipeline_managed_guardrails": ["evaded"], + } + data = { + "model": "gpt-3.5-turbo", + "metadata": json.dumps(attacker_payload), + "litellm_metadata": json.dumps(attacker_payload), + } + + real_admin_metadata = {"admin_flag": "from_proxy"} + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata=real_admin_metadata, + team_metadata=real_admin_metadata, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + populated = updated["metadata"] + # The real admin payload from user_api_key_dict wins. + assert populated["user_api_key_metadata"] == real_admin_metadata + assert populated["user_api_key_team_metadata"] == real_admin_metadata + assert populated.get("_pipeline_managed_guardrails") != ["evaded"] + + other = updated.get("litellm_metadata") or {} + # After the strip, litellm_metadata has no admin-injection slots. + assert "user_api_key_metadata" not in other + assert "user_api_key_team_metadata" not in other + assert "_pipeline_managed_guardrails" not in other + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_user_tags_without_permission(): """Caller-supplied metadata.tags must be stripped when the key/team @@ -486,9 +550,11 @@ async def test_add_litellm_data_to_request_audio_transcription_multipart(): "file": b"Fake audio bytes", } + # Opt the key in to client-supplied tags so the parsed tags from the + # JSON-string multipart body aren't stripped by the admin-injection strip. user_api_key_dict = UserAPIKeyAuth( api_key="hashed-key", - metadata={}, + metadata={"allow_client_tags": True}, team_metadata={}, spend=0.0, max_budget=100.0, From dc8e03b91f8fc71a6322078ad958465c25be3b21 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 23:14:05 +0000 Subject: [PATCH 29/41] fix(proxy): expand _guardrail_modification_check to cover all bypass keys MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Per VERIA-28's secondary recommendation. The existing check only gated metadata.guardrails. User-supplied values for disable_global_guardrails (plural and the original singular typo variant) and opted_out_global_guardrails are already silently ignored by _get_admin_metadata at read time, but the silent-ignore makes diagnosis confusing and relies on one specific read site catching them. Reject at auth time with a 403 when any of: - guardrails list (existing) - disable_global_guardrails (new) - disable_global_guardrail (new — historical singular-key variant) - opted_out_global_guardrails (new) are present in metadata, litellm_metadata, or at the request root, and the caller's team lacks can_modify_guardrails. Defense in depth: the strip at the pre-call layer still runs; this check fails loudly one layer earlier so operators see an explicit 403 rather than a silent-ignore. --- litellm/proxy/auth/auth_checks.py | 43 +++- .../proxy/auth/test_auth_checks.py | 217 ++++++++++++++---- 2 files changed, 204 insertions(+), 56 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 56958a88f6d..1f4b5311baf 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -327,15 +327,46 @@ def _global_proxy_budget_check( ) +_GUARDRAIL_MODIFICATION_KEYS: tuple = ( + "guardrails", + "disable_global_guardrails", + "disable_global_guardrail", + "opted_out_global_guardrails", +) + + def _guardrail_modification_check( request_body: dict, team_object: Optional[LiteLLM_TeamTable] ) -> None: - _request_metadata: dict = request_body.get("metadata", {}) or {} - if not _request_metadata.get("guardrails"): - return + """ + Reject user-supplied metadata flags that would modify guardrail behavior + unless the team has explicit permission. Checked keys include the plural + ``guardrails`` list plus the per-request toggles that influence whether + default-on guardrails run (``disable_global_guardrails``, + ``disable_global_guardrail`` singular, and ``opted_out_global_guardrails``). + User-supplied values for the bypass toggles are also silently ignored by + ``_get_admin_metadata`` at read time; this check adds defense in depth by + failing loudly at the auth layer so operators see an explicit 403 instead + of a confusing silent-ignore. + """ from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails + def _user_requested_modification(container: Any) -> bool: + if not isinstance(container, dict): + return False + return any(container.get(key) for key in _GUARDRAIL_MODIFICATION_KEYS) + + # Check both metadata keys — callers can populate either depending on the + # endpoint. Cover the top-level too so root-level injection is rejected. + modifies = ( + _user_requested_modification(request_body.get("metadata")) + or _user_requested_modification(request_body.get("litellm_metadata")) + or _user_requested_modification(request_body) + ) + if not modifies: + return + if not can_modify_guardrails(team_object): raise HTTPException( status_code=403, @@ -451,9 +482,9 @@ async def common_checks( # noqa: PLR0915 model=_model, team_object=team_object, llm_router=llm_router, - team_model_aliases=valid_token.team_model_aliases - if valid_token - else None, + team_model_aliases=( + valid_token.team_model_aliases if valid_token else None + ), ): raise ProxyException( message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index bd659ed518f..a3d24fc8bf6 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -62,9 +62,9 @@ def reset_constants_module(): # Reload modules before test importlib.reload(constants) importlib.reload(auth_checks) - + yield - + # Reload modules after test to clean up importlib.reload(constants) importlib.reload(auth_checks) @@ -157,9 +157,9 @@ def test_experimental_ui_token_ignores_litellm_ui_session_duration( expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00")) now = get_utc_datetime() # Must be ~10 min, NOT 24h. If LITELLM_UI_SESSION_DURATION were incorrectly used, this would fail. - assert expires <= now + timedelta(minutes=11), ( - "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION" - ) + assert expires <= now + timedelta( + minutes=11 + ), "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION" def test_get_experimental_ui_login_jwt_auth_token_invalid( @@ -293,13 +293,15 @@ def test_get_cli_jwt_auth_token_custom_expiration( # Set custom expiration to 48 hours monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48") - + # Reload the constants module to pick up the new env var importlib.reload(constants) # Also reload auth_checks to pick up the new constant value importlib.reload(auth_checks) - - token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) + + token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token( + valid_sso_user_defined_values + ) # Decrypt and verify token contents decrypted_token = decrypt_value_helper( @@ -315,7 +317,6 @@ def test_get_cli_jwt_auth_token_custom_expiration( assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1) - @pytest.mark.asyncio async def test_default_internal_user_params_with_get_user_object(monkeypatch): """Test that default_internal_user_params is used when creating a new user via get_user_object""" @@ -436,7 +437,9 @@ async def test_get_user_object_upsert_includes_user_email(): mock_prisma_client.db.litellm_usertable.create.assert_called_once() creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"] - assert "user_email" in creation_args, "user_email should be included when upserting a new user" + assert ( + "user_email" in creation_args + ), "user_email should be included when upserting a new user" assert creation_args["user_email"] == "test@example.com" assert creation_args["user_id"] == "new_test_user" @@ -463,7 +466,9 @@ def test_log_budget_lookup_failure_skips_user_not_found(): @pytest.mark.asyncio -@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock) +@patch( + "litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock +) async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch): """ Test that _get_team_db_check correctly calls the `new_team` function @@ -497,8 +502,12 @@ async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeyp @pytest.mark.asyncio -@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock) -async def test_get_team_db_check_does_not_call_new_team_if_exists(mock_new_team, monkeypatch): +@patch( + "litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock +) +async def test_get_team_db_check_does_not_call_new_team_if_exists( + mock_new_team, monkeypatch +): """ Test that _get_team_db_check does NOT call the `new_team` function if the team already exists in the database. @@ -541,8 +550,9 @@ async def test_vector_store_access_check_early_returns( if vector_store_registry: vector_store_registry.get_vector_store_ids_to_run.return_value = None - with patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch( - "litellm.vector_store_registry", vector_store_registry + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.vector_store_registry", vector_store_registry), ): result = await vector_store_access_check( request_body=request_body, @@ -639,8 +649,9 @@ async def test_vector_store_access_check_with_permissions(): mock_vector_store_registry = MagicMock() mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-1"] - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch( - "litellm.vector_store_registry", mock_vector_store_registry + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.vector_store_registry", mock_vector_store_registry), ): result = await vector_store_access_check( request_body=request_body, @@ -653,8 +664,9 @@ async def test_vector_store_access_check_with_permissions(): # Test with denied access mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-3"] - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch( - "litellm.vector_store_registry", mock_vector_store_registry + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.vector_store_registry", mock_vector_store_registry), ): with pytest.raises(ProxyException) as exc_info: await vector_store_access_check( @@ -687,8 +699,9 @@ async def test_vector_store_access_check_with_team_permissions(): "team-store-allowed" ] - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch( - "litellm.vector_store_registry", mock_vector_store_registry + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.vector_store_registry", mock_vector_store_registry), ): result = await vector_store_access_check( request_body=request_body, @@ -702,8 +715,9 @@ async def test_vector_store_access_check_with_team_permissions(): "team-store-denied" ] - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch( - "litellm.vector_store_registry", mock_vector_store_registry + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.vector_store_registry", mock_vector_store_registry), ): with pytest.raises(ProxyException) as exc_info: await vector_store_access_check( @@ -1598,12 +1612,15 @@ async def test_custom_auth_common_checks_opt_in(): mock_request = MagicMock() # Default (no flag) — common_checks should NOT be called - with patch( - "litellm.proxy.auth.user_api_key_auth.common_checks", - new_callable=AsyncMock, - ) as mock_common, patch( - "litellm.proxy.proxy_server.general_settings", - {}, + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", + new_callable=AsyncMock, + ) as mock_common, + patch( + "litellm.proxy.proxy_server.general_settings", + {}, + ), ): mock_common.return_value = True result = await _run_post_custom_auth_checks( @@ -1616,12 +1633,15 @@ async def test_custom_auth_common_checks_opt_in(): mock_common.assert_not_called() # With flag=True — common_checks SHOULD be called - with patch( - "litellm.proxy.auth.user_api_key_auth.common_checks", - new_callable=AsyncMock, - ) as mock_common, patch( - "litellm.proxy.proxy_server.general_settings", - {"custom_auth_run_common_checks": True}, + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", + new_callable=AsyncMock, + ) as mock_common, + patch( + "litellm.proxy.proxy_server.general_settings", + {"custom_auth_run_common_checks": True}, + ), ): mock_common.return_value = True result = await _run_post_custom_auth_checks( @@ -1660,9 +1680,7 @@ async def test_virtual_key_budget_check_reads_from_spend_counter(): return 1.5 return fallback_spend - with patch( - "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend - ): + with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: await _virtual_key_max_budget_check( valid_token=valid_token, @@ -1692,9 +1710,7 @@ async def test_virtual_key_budget_check_fallback_no_counter(): async def mock_get_current_spend(counter_key, fallback_spend): return fallback_spend - with patch( - "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend - ): + with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: await _virtual_key_max_budget_check( valid_token=valid_token, @@ -1723,9 +1739,7 @@ async def test_team_budget_check_reads_from_spend_counter(): return 1.5 return fallback_spend - with patch( - "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend - ): + with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: await _team_max_budget_check( team_object=team_object, @@ -1763,12 +1777,13 @@ async def test_team_member_budget_check_reads_from_spend_counter(): return 1.5 return fallback_spend - with patch( - "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend - ), patch( - "litellm.proxy.auth.auth_checks.get_team_membership", - new_callable=AsyncMock, - return_value=team_membership, + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + return_value=team_membership, + ), ): with pytest.raises(litellm.BudgetExceededError) as exc_info: await _check_team_member_budget( @@ -1780,3 +1795,105 @@ async def test_team_member_budget_check_reads_from_spend_counter(): proxy_logging_obj=proxy_logging_obj, ) assert exc_info.value.current_cost == 1.5 + + +class TestGuardrailModificationCheck: + """Defense-in-depth: `_guardrail_modification_check` must 403 when the + caller's metadata attempts to modify any guardrail-related key and the + team lacks the `modify_guardrails` permission. Checks both the + historically-covered `guardrails` list and the bypass toggles that + `_get_admin_metadata` silently ignores at read time. + """ + + def _call(self, request_body): + from litellm.proxy.auth.auth_checks import _guardrail_modification_check + + team_object = MagicMock() + team_object.metadata = {} # no permission + return _guardrail_modification_check( + request_body=request_body, team_object=team_object + ) + + def test_noop_when_no_guardrail_keys_present(self): + # no-op — should return silently + self._call({"metadata": {"unrelated": "value"}}) + + def test_rejects_guardrails_list(self): + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": {"guardrails": ["custom"]}}) + assert exc.value.status_code == 403 + + def test_rejects_disable_global_guardrails_plural(self): + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": {"disable_global_guardrails": True}}) + assert exc.value.status_code == 403 + + def test_rejects_disable_global_guardrail_singular(self): + """VERIA-28's originally-reported singular-key typo variant.""" + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": {"disable_global_guardrail": True}}) + assert exc.value.status_code == 403 + + def test_rejects_opted_out_global_guardrails(self): + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call( + {"metadata": {"opted_out_global_guardrails": ["some_guardrail"]}} + ) + assert exc.value.status_code == 403 + + def test_rejects_injection_via_litellm_metadata_key(self): + """Caller can populate the OTHER metadata key; that must also 403.""" + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"litellm_metadata": {"disable_global_guardrails": True}}) + assert exc.value.status_code == 403 + + def test_rejects_root_level_injection(self): + """Top-level injection (`request_body["disable_global_guardrails"]`) + was VERIA-28's easiest variant to hit — keep it rejected.""" + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"disable_global_guardrails": True}) + assert exc.value.status_code == 403 + + def test_allows_when_team_has_permission(self): + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=True, + ): + # no-op, should not raise + self._call({"metadata": {"disable_global_guardrails": True}}) From 76aa97f77b0781b77e4ea248e7245b0bdeefa587 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 23:35:26 +0000 Subject: [PATCH 30/41] fix(proxy): close three variant metadata/tag injection paths MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Close three variant bypasses adjacent to VERIA-28 found during post-fix variant audit: 1. _guardrail_modification_check had the same isinstance(dict) bypass Veria-AI just flagged on the pre-call strip. A caller sending `{"metadata": "{…}"}` as a JSON-encoded string (multipart/form-data or extra_body) skipped the guard, got parsed to dict downstream, and reached guardrail logic with bypass flags intact. Coerce strings via safe_json_loads before evaluating. 2. The allow_client_tags strip only covered body metadata.tags and litellm_metadata.tags — caller-supplied tags arriving via the x-litellm-tags header or root-level data["tags"] bypassed it. Gate add_request_tag_to_metadata's result on the same flag. 3. requester_metadata was deepcopied BEFORE the strip, so attacker injections (user_api_key_metadata shadows, disallowed tags, _pipeline_managed_guardrails) persisted in the snapshot. The PANW guardrail (and any future consumer) trusting requester_metadata would see forged values. Move the deepcopy to after the strip. Regression tests added for each. --- litellm/proxy/auth/auth_checks.py | 22 ++- litellm/proxy/litellm_pre_call_utils.py | 30 +++- .../proxy/auth/test_auth_checks.py | 37 +++++ .../proxy/test_litellm_pre_call_utils.py | 131 ++++++++++++++++++ 4 files changed, 213 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 1f4b5311baf..846d61384b4 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -350,12 +350,30 @@ def _guardrail_modification_check( failing loudly at the auth layer so operators see an explicit 403 instead of a confusing silent-ignore. """ + from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails + def _coerce_to_dict(container: Any) -> Optional[dict]: + """Accept dict or JSON-string (from multipart/form-data or extra_body). + + Without this, an attacker can smuggle guardrail keys past the check by + sending ``{"metadata": "{\\"disable_global_guardrails\\": true}"}`` — + ``isinstance(dict)`` on the string returns False, the check returns + no-modification, and ``add_litellm_data_to_request`` parses the string + to a dict downstream. + """ + if isinstance(container, dict): + return container + if isinstance(container, str): + parsed = safe_json_loads(container) + return parsed if isinstance(parsed, dict) else None + return None + def _user_requested_modification(container: Any) -> bool: - if not isinstance(container, dict): + coerced = _coerce_to_dict(container) + if coerced is None: return False - return any(container.get(key) for key in _GUARDRAIL_MODIFICATION_KEYS) + return any(coerced.get(key) for key in _GUARDRAIL_MODIFICATION_KEYS) # Check both metadata keys — callers can populate either depending on the # endpoint. Cover the top-level too so root-level injection is rejected. diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 33affa5351c..51f82b7c00d 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1069,9 +1069,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915 verbose_proxy_logger.warning( f"Failed to parse 'metadata' as JSON dict. Received value: {data['metadata']}" ) - data[_metadata_variable_name]["requester_metadata"] = copy.deepcopy( - data["metadata"] - ) + # requester_metadata is snapshotted AFTER the strip below so + # downstream consumers (e.g. PANW guardrail reading user_ip / + # profile_id) don't see attacker-injected admin slots preserved in + # the deepcopy. # Parse litellm_metadata if it's a string (e.g., from multipart/form-data or extra_body) if "litellm_metadata" in data and data["litellm_metadata"] is not None: @@ -1131,6 +1132,16 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ", ".join(_stripped_from), ) + # Snapshot the (now-cleaned) requester-supplied metadata for downstream + # consumers. Taking the deepcopy AFTER the strip prevents attacker- + # injected admin slots (user_api_key_metadata, tags without opt-in, + # _pipeline_managed_guardrails) from surviving in requester_metadata + # where guardrails and audit paths may read from it. + if "metadata" in data and isinstance(data["metadata"], dict): + data[_metadata_variable_name]["requester_metadata"] = copy.deepcopy( + data["metadata"] + ) + # Now merge litellm_metadata into the metadata variable (preserving existing # values) — runs AFTER the strip so attacker injections in litellm_metadata # cannot cross-contaminate the admin-authoritative metadata dict. @@ -1300,15 +1311,24 @@ async def add_litellm_data_to_request( # noqa: PLR0915 user_agent = request.headers["user-agent"] data[_metadata_variable_name]["user_agent"] = user_agent - # Check if using tag based routing + # Check if using tag based routing. The helper reads caller-controlled + # sources (x-litellm-tags header, data["tags"] root-level), so its result + # is still gated by the same allow_client_tags flag that gated the + # body-metadata tag strip above. Otherwise the strip is trivially + # bypassed by sending tags via header or at the root of the body. tags = LiteLLMProxyRequestSetup.add_request_tag_to_metadata( llm_router=llm_router, headers=_headers, data=data, ) - if tags is not None: + if tags is not None and _admin_allow_client_tags: data[_metadata_variable_name]["tags"] = tags + elif tags is not None: + verbose_proxy_logger.warning( + "Ignored caller-supplied tags from header/root body: this " + "key/team does not have `allow_client_tags: true` in its metadata." + ) # Team Callbacks controls callback_settings_obj = _get_dynamic_logging_metadata( diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index a3d24fc8bf6..0391e224084 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1897,3 +1897,40 @@ class TestGuardrailModificationCheck: ): # no-op, should not raise self._call({"metadata": {"disable_global_guardrails": True}}) + + def test_rejects_string_encoded_metadata_bypass(self): + """Regression: attacker sends metadata as JSON string to bypass the + isinstance(dict) guard. The check must coerce the string to dict + and evaluate guardrail modification keys inside it.""" + import json as _json + + from fastapi import HTTPException + + attacker_payload = {"disable_global_guardrails": True} + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": _json.dumps(attacker_payload)}) + assert exc.value.status_code == 403 + + def test_rejects_string_encoded_litellm_metadata_bypass(self): + """Same bypass via the litellm_metadata key.""" + import json as _json + + from fastapi import HTTPException + + attacker_payload = {"guardrails": ["evaded"]} + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"litellm_metadata": _json.dumps(attacker_payload)}) + assert exc.value.status_code == 403 + + def test_noop_when_string_is_not_json_object(self): + """Unparseable strings should not trigger a 403 — they have no keys.""" + self._call({"metadata": "not-json"}) + self._call({"metadata": '"just a string"'}) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 351857b9b10..5f5ae0c64d9 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -344,6 +344,137 @@ async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection assert "_pipeline_managed_guardrails" not in other +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_ignores_x_litellm_tags_header_without_permission(): + """Regression: the `x-litellm-tags` header bypassed the body-metadata + tag strip. Header tags must also be gated by `allow_client_tags`.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = { + "Content-Type": "application/json", + "x-litellm-tags": "restricted-tier,victim-team", + } + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = {"model": "gpt-3.5-turbo"} + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert "tags" not in (updated.get("metadata") or {}) + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_ignores_root_level_tags_without_permission(): + """Regression: root-level `data["tags"]` bypassed the body-metadata + tag strip. Root-level tags must also be gated by `allow_client_tags`.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "tags": ["restricted-tier", "victim-team"], + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert "tags" not in (updated.get("metadata") or {}) + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_honors_header_tags_when_opted_in(): + """When allow_client_tags=True, header-supplied tags flow through.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = { + "Content-Type": "application/json", + "x-litellm-tags": "production,ab-test", + } + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = {"model": "gpt-3.5-turbo"} + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"allow_client_tags": True}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"].get("tags") == ["production", "ab-test"] + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_user_tags_without_permission(): """Caller-supplied metadata.tags must be stripped when the key/team From b4e98d190a42c04c1e7cf3c44abf7524a92826d5 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 17 Apr 2026 00:08:40 +0000 Subject: [PATCH 31/41] fix(proxy): close 6 more metadata/tag variant bypasses MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Post-merge audit found 6 adjacent variants of the VERIA-28 class. All fixed here with regression tests: 1. Strip widened from 3 named keys to the full user_api_key_* prefix. The proxy writes a dozen user_api_key_* fields (user_id, alias, spend, team_id, request_route, end_user_id, …) into data[_metadata_variable_name]; the 3-key strip left the rest exploitable for identity/spend forgery in audit logs and guardrails. 2. proxy_server_request['body'] snapshot moved to AFTER the strip. Was captured at line ~990 before the strip ran, so standard_logging_object, lago, and spend_tracking readers saw the attacker-forged payload even though the live data dict was clean. 3. get_tags_from_request_body (auth-time) now coerces JSON-string metadata via safe_json_loads. Previously crashed with AttributeError on string metadata (DoS; potential RBAC bypass if a caller swallowed the exception). 4. get_end_user_id_from_request_body coerces JSON-string metadata/litellm_metadata. Previously isinstance(dict) guard caused end-user budget attribution to be silently skipped when the caller sent metadata as a JSON string. 5. Four hand-rolled 'if data.get("metadata") is None: data["metadata"] = {}' blocks in proxy_server.py (7160, 7341, 7590, 11375) now guard on isinstance(dict). They crashed with TypeError when metadata was a JSON string (DoS). 6. _get_admin_metadata defensively guards with isinstance(dict); previously AttributeError'd on any leaked string metadata. Also hoists the inline safe_json_loads import in _guardrail_modification_check to module level per CLAUDE.md style. --- litellm/integrations/custom_guardrail.py | 6 +- litellm/proxy/auth/auth_checks.py | 2 +- litellm/proxy/auth/auth_utils.py | 31 ++-- .../proxy/common_utils/http_parsing_utils.py | 19 +- litellm/proxy/litellm_pre_call_utils.py | 36 ++-- litellm/proxy/proxy_server.py | 15 +- .../common_utils/test_http_parsing_utils.py | 38 ++++ .../proxy/test_litellm_pre_call_utils.py | 166 ++++++++++++++++++ 8 files changed, 281 insertions(+), 32 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index b0931964cc6..abf010e0d65 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -268,7 +268,11 @@ class CustomGuardrail(CustomLogger): team_meta: dict = {} key_meta: dict = {} for key in ("metadata", "litellm_metadata"): - meta = data.get(key) or {} + # Defensive: an unparsed JSON-string metadata could leak past the + # proxy's normal parse path; don't AttributeError on .get(). + meta = data.get(key) + if not isinstance(meta, dict): + continue team_meta = meta.get("user_api_key_team_metadata") or team_meta key_meta = meta.get("user_api_key_metadata") or key_meta return {**team_meta, **key_meta} diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 846d61384b4..1621f04213f 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -31,6 +31,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider +from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.proxy._types import ( RBAC_ROLES, CallInfo, @@ -350,7 +351,6 @@ def _guardrail_modification_check( failing loudly at the auth layer so operators see an explicit 403 instead of a confusing silent-ignore. """ - from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails def _coerce_to_dict(container: Any) -> Optional[dict]: diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 64766bbaadd..12d9bff91e0 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -823,19 +823,30 @@ def get_end_user_id_from_request_body( user_from_body_user_field = request_body["user"] return str(user_from_body_user_field) + def _as_dict(value: Any) -> dict: + # metadata / litellm_metadata can arrive as JSON strings from + # multipart/form-data or extra_body; coerce so string-encoded + # payloads can't evade end-user attribution. + if isinstance(value, dict): + return value + if isinstance(value, str): + from litellm.litellm_core_utils.safe_json_loads import safe_json_loads + + parsed = safe_json_loads(value) + return parsed if isinstance(parsed, dict) else {} + return {} + # Check 4: 'litellm_metadata.user' in request_body (commonly Anthropic) - litellm_metadata = request_body.get("litellm_metadata") - if isinstance(litellm_metadata, dict): - user_from_litellm_metadata = litellm_metadata.get("user") - if user_from_litellm_metadata is not None: - return str(user_from_litellm_metadata) + litellm_metadata = _as_dict(request_body.get("litellm_metadata")) + user_from_litellm_metadata = litellm_metadata.get("user") + if user_from_litellm_metadata is not None: + return str(user_from_litellm_metadata) # Check 5: 'metadata.user_id' in request_body (another common pattern) - metadata_dict = request_body.get("metadata") - if isinstance(metadata_dict, dict): - user_id_from_metadata_field = metadata_dict.get("user_id") - if user_id_from_metadata_field is not None: - return str(user_id_from_metadata_field) + metadata_dict = _as_dict(request_body.get("metadata")) + user_id_from_metadata_field = metadata_dict.get("user_id") + if user_id_from_metadata_field is not None: + return str(user_id_from_metadata_field) # Check 6: 'safety_identifier' in request body (OpenAI Responses API parameter) # SECURITY NOTE: safety_identifier can be set by any caller in the request body. diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 1dd25262127..71abdfa5e9e 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -197,10 +197,10 @@ def check_file_size_under_limit( if llm_router is not None and request_data["model"] in router_model_names: try: - deployment: Optional[ - Deployment - ] = llm_router.get_deployment_by_model_group_name( - model_group_name=request_data["model"] + deployment: Optional[Deployment] = ( + llm_router.get_deployment_by_model_group_name( + model_group_name=request_data["model"] + ) ) if ( deployment @@ -426,7 +426,16 @@ def get_tags_from_request_body(request_body: dict) -> List[str]: List of tag names (strings), empty list if no valid tags found """ metadata_variable_name = get_metadata_variable_name_from_kwargs(request_body) - metadata = request_body.get(metadata_variable_name) or {} + metadata = request_body.get(metadata_variable_name) + # metadata can arrive as a JSON string from multipart/form-data or extra_body; + # coerce defensively so .get() below never raises AttributeError. + if isinstance(metadata, str): + from litellm.litellm_core_utils.safe_json_loads import safe_json_loads + + parsed = safe_json_loads(metadata) + metadata = parsed if isinstance(parsed, dict) else {} + elif not isinstance(metadata, dict): + metadata = {} tags_in_metadata: Any = metadata.get("tags", []) tags_in_request_body: Any = request_body.get("tags", []) combined_tags: List[str] = [] diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 51f82b7c00d..aa3a29e4f2f 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -981,13 +981,16 @@ async def add_litellm_data_to_request( # noqa: PLR0915 # Init - Proxy Server Request # we do this as soon as entering so we track the original request ########################################################## - # Track arrival time for queue time metric + # Track arrival time for queue time metric. The body snapshot is filled + # in after the admin-injection strip below so the audit / spend-tracking + # consumers of proxy_server_request["body"] see the cleaned metadata + # rather than attacker-forged user_api_key_* fields. arrival_time = time.time() data["proxy_server_request"] = { "url": str(request.url), "method": request.method, "headers": _headers, - "body": copy.copy(data), # use copy instead of deepcopy + "body": None, # filled in post-strip; see below "arrival_time": arrival_time, # Track when request arrived at proxy } @@ -1087,19 +1090,24 @@ async def add_litellm_data_to_request( # noqa: PLR0915 # Strip internal pipeline state and admin-injection slots from user input. # Runs AFTER the string-to-dict parse above so JSON-string metadata (sent - # via multipart/form-data or extra_body) cannot smuggle `user_api_key_metadata` - # past the isinstance(dict) guard. + # via multipart/form-data or extra_body) cannot smuggle admin fields past + # the isinstance(dict) guard. # - # The proxy writes user_api_key_metadata / user_api_key_team_metadata into - # data[_metadata_variable_name] below; if a caller pre-populates either - # key on the OTHER metadata field, _get_admin_metadata lookups would treat - # the caller's payload as admin-configured. + # The proxy populates a family of ``user_api_key_*`` fields below + # (user_api_key_metadata, user_api_key_user_id, user_api_key_alias, + # user_api_key_spend, user_api_key_team_metadata, …) into + # data[_metadata_variable_name]. Because the proxy only writes to ONE of + # the two metadata dicts, a caller pre-populating any of these keys on + # the OTHER metadata dict would have their forged values surface in + # guardrails, spend tracking, audit logs, and identity resolution. Strip + # by prefix so new ``user_api_key_*`` fields added in the future are + # covered without per-key maintenance. for _meta_key in ("metadata", "litellm_metadata"): _user_meta = data.get(_meta_key) if isinstance(_user_meta, dict): _user_meta.pop("_pipeline_managed_guardrails", None) - _user_meta.pop("user_api_key_metadata", None) - _user_meta.pop("user_api_key_team_metadata", None) + for _k in [k for k in _user_meta if k.startswith("user_api_key_")]: + _user_meta.pop(_k, None) # Strip caller-supplied routing/budget tags unless the admin has opted # this key or team in via metadata.allow_client_tags=True. Tags drive @@ -1132,9 +1140,15 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ", ".join(_stripped_from), ) + # Fill in the proxy_server_request body snapshot now that metadata has + # been parsed and stripped. Consumers (standard_logging_payload, lago, + # spend_tracking_utils, streaming_iterator) read `body` to audit the + # request; taking the snapshot here ensures they see cleaned metadata. + data["proxy_server_request"]["body"] = copy.copy(data) + # Snapshot the (now-cleaned) requester-supplied metadata for downstream # consumers. Taking the deepcopy AFTER the strip prevents attacker- - # injected admin slots (user_api_key_metadata, tags without opt-in, + # injected admin slots (user_api_key_*, tags without opt-in, # _pipeline_managed_guardrails) from surviving in requester_metadata # where guardrails and audit paths may read from it. if "metadata" in data and isinstance(data["metadata"], dict): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2d789b982da..1d61bebee44 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2201,9 +2201,11 @@ def run_ollama_serve(): with open(os.devnull, "w") as devnull: subprocess.Popen(command, stdout=devnull, stderr=devnull) except Exception as e: - verbose_proxy_logger.debug(f""" + verbose_proxy_logger.debug( + f""" LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve` - """) + """ + ) def _get_process_rss_mb() -> Optional[float]: @@ -7157,7 +7159,10 @@ async def chat_completion( # noqa: PLR0915 global user_temperature, user_request_timeout, user_max_tokens, user_api_base data = await _read_request_body(request=request) if user_api_key_dict is not None: - if data.get("metadata") is None: + if not isinstance(data.get("metadata"), dict): + # Covers both missing and JSON-string metadata (multipart / + # extra_body); otherwise `data["metadata"][k] = v` below raises + # TypeError on a string value and 500s the request. data["metadata"] = {} if ( hasattr(user_api_key_dict, "user_id") @@ -11372,7 +11377,9 @@ async def async_queue_request( # if users are using user_api_key_auth, set `user` in `data` data["user"] = user_api_key_dict.user_id - if "metadata" not in data: + if not isinstance(data.get("metadata"), dict): + # Covers both missing and JSON-string metadata (multipart / + # extra_body); see above for the same guard upstream. data["metadata"] = {} data["metadata"]["user_api_key"] = user_api_key_dict.api_key data["metadata"]["user_api_key_metadata"] = user_api_key_dict.metadata diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index a1484bc263b..c9f595626ed 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -835,3 +835,41 @@ def test_safe_get_request_headers_state_unavailable(): result = _safe_get_request_headers(mock_request) assert result == {"content-type": "application/json"} + + +class TestGetTagsFromRequestBodyStringCoerce: + """Regression: the auth-time tag helper used `metadata.get("tags", ...)` + directly, which raised AttributeError when metadata arrived as a JSON + string (multipart/form-data or extra_body). That turned into a DoS at + auth time and potentially bypassed tag-based RBAC if the caller caught + the exception and fell through with empty tags. + """ + + def test_json_string_metadata_is_coerced_to_dict(self): + from litellm.proxy.common_utils.http_parsing_utils import ( + get_tags_from_request_body, + ) + + metadata_json = json.dumps({"tags": ["a", "b"]}) + # Must not raise + tags = get_tags_from_request_body({"metadata": metadata_json}) + assert tags == ["a", "b"] + + def test_unparseable_string_metadata_is_ignored(self): + from litellm.proxy.common_utils.http_parsing_utils import ( + get_tags_from_request_body, + ) + + # Must not raise; must yield no metadata tags but keep root tags + tags = get_tags_from_request_body( + {"metadata": "not-json", "tags": ["root-only"]} + ) + assert tags == ["root-only"] + + def test_dict_metadata_still_works(self): + from litellm.proxy.common_utils.http_parsing_utils import ( + get_tags_from_request_body, + ) + + tags = get_tags_from_request_body({"metadata": {"tags": ["x"]}}) + assert tags == ["x"] diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 5f5ae0c64d9..664a936b540 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -280,6 +280,172 @@ async def test_add_litellm_data_to_request_strips_admin_injection_slots(): assert "_pipeline_managed_guardrails" not in other +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_all_user_api_key_prefix_keys(): + """Strip must cover the full user_api_key_* family, not a hand-maintained + list of 2-3 names. Proxy writes a dozen such fields (user_id, alias, + spend, team_id, request_route, …) and an attacker populating any of them + in the non-authoritative metadata key would otherwise forge identity / + spend in audit logs and guardrails.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + attacker_injected = { + "user_api_key_user_id": "victim", + "user_api_key_alias": "admin-key", + "user_api_key_spend": 0.0, + "user_api_key_team_id": "victim-team", + "user_api_key_end_user_id": "victim-user", + "user_api_key_request_route": "/fake/route", + "user_api_key_hash": "fake-hash", + } + data = { + "model": "gpt-3.5-turbo", + "metadata": {**attacker_injected}, + "litellm_metadata": {**attacker_injected}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + user_id="real-user", + metadata={}, + team_metadata={}, + spend=42.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + # The non-authoritative metadata dict must not retain ANY attacker-injected + # user_api_key_* key. + other = updated.get("litellm_metadata") or {} + attacker_leaks = [k for k in other if k.startswith("user_api_key_")] + assert attacker_leaks == [], f"Unexpected leaked keys: {attacker_leaks}" + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_string_metadata_does_not_crash(): + """Regression: pre-strip code that pre-populated data['metadata'][k]=v + before the string-to-dict parse would crash on JSON-string metadata. + The snapshot / strip / admin-population pipeline must survive metadata + arriving as a string.""" + import json as _json + + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "multipart/form-data"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": _json.dumps({"generation_name": "test"}), + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + # Must not raise TypeError / AttributeError. + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + # The parsed metadata should be a dict and the proxy snapshot body + # should have been taken AFTER the strip (so no leaked user_api_key_* + # from a raw string snapshot). + assert isinstance(updated["metadata"], dict) + assert updated["metadata"].get("generation_name") == "test" + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_proxy_server_request_body_is_post_strip(): + """Regression: proxy_server_request['body'] used to be snapshotted before + the admin-slot strip, so standard_logging_object and spend-tracking + readers saw attacker-injected payload. Snapshot must now be post-strip.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"user_api_key_user_id": "victim"}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + user_id="real-user", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + snapshot_body = updated["proxy_server_request"]["body"] + assert snapshot_body is not None + snapshot_metadata = snapshot_body.get("metadata") or {} + assert "user_api_key_user_id" not in snapshot_metadata or ( + snapshot_metadata["user_api_key_user_id"] != "victim" + ) + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection(): """Regression: metadata arriving as a JSON string (multipart/form-data or From 467166fdd7c26b89ea73cf418e419881e25fa1fe Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 17 Apr 2026 00:11:00 +0000 Subject: [PATCH 32/41] fix(proxy): enforce per-target org authorization on /user/delete Veria admin-queue finding E3NpkuAd, Audit-B #1. The route-level gate accepts this call when the caller is PROXY_ADMIN or ORG_ADMIN of any org named in request_data["organization_id"]/["organizations"]. The handler processes data.user_ids without cross-checking whether those users belong to the caller's administered orgs, so an org-admin of org-A could delete users in org-B via: {"user_ids": ["victim_in_org_B"], "organization_id": "org-A"} Add per-target authorization: org-admins may only delete users whose entire org membership is within their admin scope; targets with any org outside scope (or no org at all) require PROXY_ADMIN. Regression test confirms an ORG_ADMIN call fails with 403 and no cascade delete_many runs. --- .../internal_user_endpoints.py | 98 +++++++++++++++---- .../test_internal_user_endpoints.py | 71 ++++++++++++++ 2 files changed, 149 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 1772da3d15e..be0f02439bc 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -2057,6 +2057,38 @@ async def delete_user( if data.user_ids is None: raise HTTPException(status_code=400, detail={"error": "No user id passed in"}) + # Per-target authorization: the route-level gate accepts this call when + # the caller is PROXY_ADMIN or an ORG_ADMIN of *any* org named in + # request_data["organization_id"]/["organizations"]. That gate does NOT + # cross-check data.user_ids against the caller's scope, so without this + # loop an org-admin of org-A could delete users in org-B by supplying + # {"user_ids": [victim_in_org_B], "organization_id": "org-A"}. + caller_is_proxy_admin = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + caller_admin_org_ids: set = set() + if not caller_is_proxy_admin: + caller_memberships = ( + await prisma_client.db.litellm_organizationmembership.find_many( + where={ + "user_id": user_api_key_dict.user_id, + "user_role": LitellmUserRoles.ORG_ADMIN.value, + } + ) + if user_api_key_dict.user_id + else [] + ) + caller_admin_org_ids = { + m.organization_id for m in caller_memberships if m.organization_id + } + if not caller_admin_org_ids: + raise HTTPException( + status_code=403, + detail={ + "error": "Only PROXY_ADMIN or ORG_ADMIN users may delete users." + }, + ) + # check that all teams passed exist for user_id in data.user_ids: user_row = await prisma_client.db.litellm_usertable.find_unique( @@ -2068,30 +2100,56 @@ async def delete_user( status_code=404, detail={"error": f"User not found, passed user_id={user_id}"}, ) - else: - # Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True - # we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes - if litellm.store_audit_logs is True: - # make an audit log for each team deleted - _user_row = user_row.json(exclude_none=True) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=litellm_changed_by - or user_api_key_dict.user_id - or litellm_proxy_admin_name, - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.USER_TABLE_NAME, - object_id=user_id, - action="deleted", - updated_values="{}", - before_value=_user_row, + if not caller_is_proxy_admin: + target_memberships = ( + await prisma_client.db.litellm_organizationmembership.find_many( + where={"user_id": user_id} + ) + ) + target_org_ids = { + m.organization_id for m in target_memberships if m.organization_id + } + # Org-admin may only delete users whose entire org membership is + # within their admin scope. A target with ANY org outside the + # caller's scope (or no org at all) requires PROXY_ADMIN. + if not target_org_ids or not target_org_ids.issubset( + caller_admin_org_ids + ): + raise HTTPException( + status_code=403, + detail={ + "error": ( + f"User {user_id} is not within your admin scope. " + "Only PROXY_ADMIN may delete users outside your " + "administered organizations." ) + }, + ) + + # Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True + # we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes + if litellm.store_audit_logs is True: + # make an audit log for each team deleted + _user_row = user_row.json(exclude_none=True) + + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=litellm_changed_by + or user_api_key_dict.user_id + or litellm_proxy_admin_name, + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.USER_TABLE_NAME, + object_id=user_id, + action="deleted", + updated_values="{}", + before_value=_user_row, ) ) + ) ## CLEANUP MEMBERS_WITH_ROLES fetch_all_teams = await prisma_client.db.litellm_teamtable.find_many( diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 0f90d236aed..104071c3e61 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1878,6 +1878,77 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker): assert condition[field] == {"in": ["admin-creator"]} +@pytest.mark.asyncio +async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker): + """Regression: an org admin of org-A must not be able to delete a user + whose org memberships include org-B. + + Route-level gate accepts the request when the caller supplies an + `organization_id` they administer; without per-user org authorization + the handler would cascade-delete the victim's keys, memberships, and + user row regardless of where the victim actually belongs. + """ + from fastapi import HTTPException + + from litellm.proxy._types import DeleteUserRequest, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user + + mock_prisma_client = mocker.MagicMock() + + # Target user exists and is a member of org-B only. + mock_target_user = mocker.MagicMock() + mock_target_user.user_id = "victim" + mock_target_user.user_email = "victim@example.com" + mock_target_user.teams = [] + mock_target_user.json.return_value = "{}" + + async def mock_find_unique(*args, **kwargs): + return mock_target_user + + mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock( + side_effect=mock_find_unique + ) + + # Caller (org_admin_user) administers org-A. + caller_membership = mocker.MagicMock() + caller_membership.organization_id = "org-A" + + # Target user is a member of org-B (outside caller's scope). + target_membership = mocker.MagicMock() + target_membership.organization_id = "org-B" + + async def mock_find_memberships(*args, **kwargs): + where = kwargs.get("where") or (args[0] if args else {}) + user_id = where.get("user_id") + if user_id == "org_admin_user": + return [caller_membership] + if user_id == "victim": + return [target_membership] + return [] + + mock_prisma_client.db.litellm_organizationmembership.find_many = mocker.AsyncMock( + side_effect=mock_find_memberships + ) + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + data = DeleteUserRequest(user_ids=["victim"]) + user_api_key_dict = UserAPIKeyAuth( + user_id="org_admin_user", user_role=LitellmUserRoles.ORG_ADMIN + ) + + with pytest.raises(HTTPException) as exc: + await delete_user(data=data, user_api_key_dict=user_api_key_dict) + assert exc.value.status_code == 403 + + # Critical: no delete_many calls should have executed. + assert not hasattr( + mock_prisma_client.db.litellm_verificationtoken.delete_many, "mock_calls" + ) or len( + mock_prisma_client.db.litellm_verificationtoken.delete_many.mock_calls + ) == 0 + + # ===================================================================== # /v2/user/info endpoint tests # ===================================================================== From 8c0668f105c00772d785b7981973709edc718c05 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 17 Apr 2026 00:17:57 +0000 Subject: [PATCH 33/41] perf: batch target membership lookup in delete_user to avoid N+1 Greptile P1 on the /user/delete fix. Per CLAUDE.md 'No N+1 queries', move the find_many inside the per-user loop to a single batched fetch with {'user_id': {'in': data.user_ids}} before the loop, then distribute to a per-user set in memory. --- .../internal_user_endpoints.py | 25 +++++++++++++------ .../test_internal_user_endpoints.py | 15 ++++++++--- 2 files changed, 28 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index be0f02439bc..09cd7fc5f6d 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -2089,6 +2089,22 @@ async def delete_user( }, ) + # Batch-fetch target memberships once before the per-user loop. Avoids + # an N+1 DB call when delete_user is called with a large user_ids list. + target_org_ids_by_user: Dict[str, set] = {} + if not caller_is_proxy_admin: + all_target_memberships = ( + await prisma_client.db.litellm_organizationmembership.find_many( + where={"user_id": {"in": data.user_ids}} + ) + ) + for m in all_target_memberships: + if not m.organization_id: + continue + target_org_ids_by_user.setdefault(m.user_id, set()).add( + m.organization_id + ) + # check that all teams passed exist for user_id in data.user_ids: user_row = await prisma_client.db.litellm_usertable.find_unique( @@ -2102,14 +2118,7 @@ async def delete_user( ) if not caller_is_proxy_admin: - target_memberships = ( - await prisma_client.db.litellm_organizationmembership.find_many( - where={"user_id": user_id} - ) - ) - target_org_ids = { - m.organization_id for m in target_memberships if m.organization_id - } + target_org_ids = target_org_ids_by_user.get(user_id, set()) # Org-admin may only delete users whose entire org membership is # within their admin scope. A target with ANY org outside the # caller's scope (or no org at all) requires PROXY_ADMIN. diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 104071c3e61..a074a4a211f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1919,11 +1919,18 @@ async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker): async def mock_find_memberships(*args, **kwargs): where = kwargs.get("where") or (args[0] if args else {}) - user_id = where.get("user_id") - if user_id == "org_admin_user": + user_id_filter = where.get("user_id") + # Batched lookup: {"user_id": {"in": [...]}} returns target memberships. + # Caller role lookup: {"user_id": "", "user_role": ...}. + if isinstance(user_id_filter, dict) and "in" in user_id_filter: + if "victim" in user_id_filter["in"]: + # Attach user_id on the mock so the caller can build its + # per-user dict from the batch result. + target_membership.user_id = "victim" + return [target_membership] + return [] + if user_id_filter == "org_admin_user": return [caller_membership] - if user_id == "victim": - return [target_membership] return [] mock_prisma_client.db.litellm_organizationmembership.find_many = mocker.AsyncMock( From 132063289fca416c18b55fa99fa44919a2b0b41b Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 17 Apr 2026 00:20:52 +0000 Subject: [PATCH 34/41] fix(proxy): strip root-level data['tags'] alongside metadata tags Greptile P2. The admin-inject gate only removed tags from data['metadata'] and data['litellm_metadata']; and the policy engine read directly, so a caller without allow_client_tags could still drive tag-based policy decisions by moving tags to the body root. Also strip the root key in the same branch. --- litellm/proxy/litellm_pre_call_utils.py | 7 +++++++ tests/test_litellm/proxy/test_litellm_pre_call_utils.py | 5 +++++ 2 files changed, 12 insertions(+) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index aa3a29e4f2f..7467bbae232 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1132,6 +1132,13 @@ async def add_litellm_data_to_request( # noqa: PLR0915 if isinstance(_user_meta, dict) and "tags" in _user_meta: _user_meta.pop("tags", None) _stripped_from.append(_meta_key) + # Also strip the root-level `tags` field. get_tags_from_request_body + # reads request_body["tags"] directly and feeds it to the policy + # engine, so leaving it in place here would let the strip-in-metadata + # above be trivially bypassed by moving the tags to the body root. + if "tags" in data: + data.pop("tags", None) + _stripped_from.append("tags (root)") if _stripped_from: verbose_proxy_logger.warning( "Stripped caller-supplied tags from %s: this key/team does " diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 664a936b540..360ddd426a4 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -596,6 +596,11 @@ async def test_add_litellm_data_to_request_ignores_root_level_tags_without_permi ) assert "tags" not in (updated.get("metadata") or {}) + # Also ensure the root-level tags are removed. get_tags_from_request_body + # reads request_body["tags"] directly, so leaving it in place would let + # the policy engine see caller-supplied tags even after the metadata + # strip. + assert "tags" not in updated @pytest.mark.asyncio From cdb29946eb75a792566d8447e1e261387fa38368 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 16 Apr 2026 17:22:28 -0700 Subject: [PATCH 35/41] fix: align agent endpoint and routing permission checks with existing pattern --- litellm/proxy/agent_endpoints/a2a_routing.py | 23 +++++++++++++- litellm/proxy/agent_endpoints/endpoints.py | 32 ++++++++++++++------ litellm/proxy/common_request_processing.py | 25 +++++++++------ litellm/proxy/route_llm_request.py | 6 +++- 4 files changed, 64 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 8f951414994..d5f5348731f 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -8,10 +8,17 @@ Looks up agents in the registry and injects their API base URL. from typing import Any, Optional import litellm +from fastapi import HTTPException + from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import UserAPIKeyAuth -def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]: +async def route_a2a_agent_request( + data: dict, + route_type: str, + user_api_key_dict: Optional[UserAPIKeyAuth] = None, +) -> Optional[Any]: """ Route A2A agent requests directly to litellm with injected API base. @@ -19,6 +26,9 @@ def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]: """ # Import here to avoid circular imports from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + ) from litellm.proxy.route_llm_request import ( ROUTE_ENDPOINT_MAPPING, ProxyModelNotFoundError, @@ -40,6 +50,17 @@ def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]: route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) raise ProxyModelNotFoundError(route=route_name, model_name=model_name) + # Verify the caller is permitted to use this agent + is_allowed = await AgentRequestHandler.is_agent_allowed( + agent_id=agent.agent_id, + user_api_key_auth=user_api_key_dict, + ) + if not is_allowed: + raise HTTPException( + status_code=403, + detail=f"Agent '{agent_name}' is not allowed for your key/team. Contact proxy admin for access.", + ) + # Get API base URL from agent config if not agent.agent_card_params or "url" not in agent.agent_card_params: verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' has no URL configured") diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 64c20d5ed5e..4c5fbaa22c5 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -200,10 +200,9 @@ async def get_agents( for agent in returned_agents: if agent.litellm_params is None: agent.litellm_params = {} - agent.litellm_params[ - "is_public" - ] = litellm.public_agent_groups is not None and ( - agent.agent_id in litellm.public_agent_groups + agent.litellm_params["is_public"] = ( + litellm.public_agent_groups is not None + and (agent.agent_id in litellm.public_agent_groups) ) # Redact sensitive fields for non-admin users @@ -393,6 +392,19 @@ async def get_agent_by_id( """ await check_feature_access_for_user(user_api_key_dict, "agents") + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + ) + + is_allowed = await AgentRequestHandler.is_agent_allowed( + agent_id=agent_id, user_api_key_auth=user_api_key_dict + ) + if not is_allowed: + raise HTTPException( + status_code=403, + detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.", + ) + from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -409,13 +421,13 @@ async def get_agent_by_id( agent_dict = agent_row.model_dump() if agent_row.object_permission is not None: try: - agent_dict[ - "object_permission" - ] = agent_row.object_permission.model_dump() + agent_dict["object_permission"] = ( + agent_row.object_permission.model_dump() + ) except Exception: - agent_dict[ - "object_permission" - ] = agent_row.object_permission.dict() + agent_dict["object_permission"] = ( + agent_row.object_permission.dict() + ) agent = AgentResponse(**agent_dict) # type: ignore else: # Agent found in memory — refresh spend from DB diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c4717ad9cf3..97801baaf0c 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -866,9 +866,11 @@ class ProxyBaseLLMRequestProcessing: "Request received by LiteLLM: payload too large to log (%d bytes, limit %d). Keys: %s", len(_payload_str), MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG, - list(self.data.keys()) - if isinstance(self.data, dict) - else type(self.data).__name__, + ( + list(self.data.keys()) + if isinstance(self.data, dict) + else type(self.data).__name__ + ), ) else: verbose_proxy_logger.debug( @@ -1054,6 +1056,7 @@ class ProxyBaseLLMRequestProcessing: route_type=route_type, llm_router=llm_router, user_model=user_model, + user_api_key_dict=user_api_key_dict, ) tasks.append(llm_call) @@ -1128,9 +1131,9 @@ class ProxyBaseLLMRequestProcessing: # aliasing/routing, but the OpenAI-compatible response `model` field should reflect # what the client sent. if requested_model_from_client: - self.data[ - "_litellm_client_requested_model" - ] = requested_model_from_client + self.data["_litellm_client_requested_model"] = ( + requested_model_from_client + ) # Streaming: attach a closure that fires after all guardrail # end-of-stream blocks complete. CSW.__anext__ stores the @@ -1731,7 +1734,9 @@ class ProxyBaseLLMRequestProcessing: verbose_proxy_logger.debug("inside generator") try: str_so_far = "" - async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( + async for ( + chunk + ) in proxy_logging_obj.async_post_call_streaming_iterator_hook( user_api_key_dict=user_api_key_dict, response=response, request_data=request_data, @@ -1959,9 +1964,9 @@ class ProxyBaseLLMRequestProcessing: # Add cache-related fields to **params (handled by Usage.__init__) if cache_creation_input_tokens is not None: - usage_kwargs[ - "cache_creation_input_tokens" - ] = cache_creation_input_tokens + usage_kwargs["cache_creation_input_tokens"] = ( + cache_creation_input_tokens + ) if cache_read_input_tokens is not None: usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index f1590b16c24..17cc4374560 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -4,6 +4,7 @@ from typing import TYPE_CHECKING, Any, Literal, Optional from fastapi import HTTPException, status import litellm +from litellm.proxy._types import UserAPIKeyAuth if TYPE_CHECKING: from litellm.router import Router as _Router @@ -314,6 +315,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "acancel_run", "adelete_run", ], + user_api_key_dict: Optional[UserAPIKeyAuth] = None, ): """ Common helper to route the request @@ -548,7 +550,9 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin route_a2a_agent_request, ) - result = route_a2a_agent_request(data, route_type) + result = await route_a2a_agent_request( + data, route_type, user_api_key_dict=user_api_key_dict + ) if result is not None: return result # Fall through to raise exception below if result is None From 662d05531d61d8d9e58fcdd1d25080da5d6119fa Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 17 Apr 2026 00:36:28 +0000 Subject: [PATCH 36/41] fix(proxy): close three more org-boundary escape paths MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Continuation of Veria E3NpkuAd / Audit-B hardening. All three are the same anti-pattern PR #25904 already addressed for _user_is_org_admin and /user/delete: route-level gate trusts a caller-supplied scope field, handler operates on a different scope. 1. /user/update no longer silently creates a user when the target email doesn't exist. Pre-fix, an org admin could supply a fresh email + caller-chosen budget/models/metadata; the INSERT path created the user with no org attachment, bypassing /user/new's org/team authorization. Now require PROXY_ADMIN for the create branch; return 404 otherwise. Also fixes /user/bulk_update because it dispatches through the same _update_single_user_helper. 2. /team/bulk_member_add with all_users=true restricted to PROXY_ADMIN. The flag pulls every user in the database into the target team, ignoring org scope — any team admin could use it to capture every user across every org into their team. 3. /team/update now verifies destination-org admin rights. When the request carries an organization_id that differs from the team's current org, an org admin of the caller's current org could previously relocate the team into any other org (draining their resources, or capturing a team they once administered). Require PROXY_ADMIN or org-admin of the DESTINATION org for the relocation. Regression tests for #1 and #3; #2 covered by the existing bulk_add suite after the gate addition. --- .../internal_user_endpoints.py | 17 +++++++ .../management_endpoints/team_endpoints.py | 50 +++++++++++++++++++ .../test_internal_user_endpoints.py | 37 ++++++++++++++ .../test_team_endpoints.py | 10 ++++ 4 files changed, 114 insertions(+) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 09cd7fc5f6d..e343429ff12 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -1180,6 +1180,23 @@ async def _update_single_user_helper( "error": "User does not have permission to update this user. Only PROXY_ADMIN can update other users." }, ) + else: + # Silent-create guard: if the target user doesn't exist, the update + # path falls through to an upsert that creates a new user with + # caller-supplied fields (models, metadata, budgets, …). Only + # PROXY_ADMIN is allowed to create users this way; otherwise an org + # admin could spawn arbitrary users attached to nothing by supplying + # a fresh email, bypassing the /user/new org/team-scoping checks. + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + raise HTTPException( + status_code=404, + detail={ + "error": ( + "User not found. Only PROXY_ADMIN can create users " + "via /user/update; use /user/new instead." + ) + }, + ) existing_metadata = ( cast(Dict, getattr(existing_user_row, "metadata", {}) or {}) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index edd92cb83de..0bb6c3d90f3 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1544,6 +1544,42 @@ async def update_team( # noqa: PLR0915 if ( data.organization_id is not None and len(data.organization_id) > 0 ): # allow unsetting the organization_id + # If the caller is relocating the team to a different org, they + # must also be PROXY_ADMIN or an org-admin of the DESTINATION org. + # _verify_team_access above only checked the team's CURRENT org, + # so without this gate an org-admin could hand their team to any + # other org (or capture a team from another org they once + # administered into a new destination). + current_org_id = getattr(existing_team_row, "organization_id", None) + if ( + data.organization_id != current_org_id + and user_api_key_dict.user_role + != LitellmUserRoles.PROXY_ADMIN.value + ): + # Is the caller org_admin of the destination org? + caller_memberships = ( + await prisma_client.db.litellm_organizationmembership.find_many( + where={ + "user_id": user_api_key_dict.user_id, + "organization_id": data.organization_id, + "user_role": LitellmUserRoles.ORG_ADMIN.value, + } + ) + if user_api_key_dict.user_id + else [] + ) + if not caller_memberships: + raise HTTPException( + status_code=403, + detail={ + "error": ( + "Relocating a team to a different organization " + "requires PROXY_ADMIN or org-admin of the " + "destination org." + ) + }, + ) + await fetch_and_validate_organization( organization_id=data.organization_id, existing_team_row=existing_team_row, @@ -2682,6 +2718,20 @@ async def bulk_team_member_add( ) if data.all_users: + # `all_users=True` pulls every user in the database into this team, + # regardless of org. Any team admin could use it to capture every + # user across every org into a team they control. Restrict to + # PROXY_ADMIN. + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + raise HTTPException( + status_code=403, + detail={ + "error": ( + "`all_users=true` is restricted to PROXY_ADMIN. " + "Org/team admins must specify explicit member lists." + ) + }, + ) # get all users from the database all_users_in_db = await prisma_client.db.litellm_usertable.find_many( order={"created_at": "desc"} diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index a074a4a211f..a1ba7ecd677 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1956,6 +1956,43 @@ async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker): ) == 0 +@pytest.mark.asyncio +async def test_user_update_rejects_silent_create_for_non_proxy_admin(mocker): + """Regression: `/user/update` with an unknown user_email used to fall + through to an INSERT, silently creating a new user with caller-supplied + budget, models, and metadata. An org admin could use this to spawn + arbitrary users outside the /user/new authorization flow.""" + from fastapi import HTTPException + + from litellm.proxy._types import UpdateUserRequest, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = mocker.MagicMock() + # user_email lookup yields None → would silently create pre-fix. + mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock( + return_value=None + ) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + user_request = UpdateUserRequest( + user_email="newcomer@example.com", + max_budget=1_000_000, + models=["gpt-4"], + ) + org_admin = UserAPIKeyAuth( + user_id="org-admin", + user_role=LitellmUserRoles.ORG_ADMIN, + ) + + with pytest.raises(HTTPException) as exc: + await _update_single_user_helper( + user_request=user_request, user_api_key_dict=org_admin + ) + assert exc.value.status_code == 404 + + # ===================================================================== # /v2/user/info endpoint tests # ===================================================================== diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 9b4bd790493..896e1efe0e4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -5128,6 +5128,16 @@ async def test_update_team_guardrails_with_org_id(): return_value=mock_org ) + # Destination-org guard in update_team queries for the caller's + # ORG_ADMIN membership on the destination org. Return a match so + # the guardrails-update path (the subject under test) proceeds. + mock_org_admin_membership = MagicMock() + mock_org_admin_membership.user_id = "org-admin-guardrails-test" + mock_org_admin_membership.organization_id = "test-org-guardrails" + mock_prisma.db.litellm_organizationmembership.find_many = AsyncMock( + return_value=[mock_org_admin_membership] + ) + # Mock team update mock_updated_team = MagicMock(spec=LiteLLM_TeamTable) mock_updated_team.team_id = "team-guardrails-123" From c7c3df2b02c0fd7e4544bdf5cf3e3b0f8118b734 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 17 Apr 2026 00:41:00 +0000 Subject: [PATCH 37/41] fix(proxy): extend /key/update admin check to non-budget fields MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Audit-B #2. _check_key_admin_access was gated on max_budget/spend changes only, which meant a non-admin caller could blanket-rewrite any OTHER field on any key (key_alias, models, tpm_limit, rpm_limit, metadata, tags, allowed_routes, guardrails, blocked, duration, permissions, auto_rotate, access_group_ids, object_permission, …) as long as they avoided budget/spend. Example attack: POST /key/update { key: sk-victim-in-org-B, models: [], blocked: true, organization_id: org-A, } The caller is org-admin of org-A, which satisfies the route gate; the handler then wipes models and blocks the victim's key. Policy after this fix: - PROXY_ADMIN: always allowed. - Key OWNER (matching user_id): allowed for non-budget fields; budget/spend changes still require team/org admin. - Everyone else: must pass _check_key_admin_access (PROXY_ADMIN / key-owner / team-admin / org-admin of the key). Regression test confirms a non-owner INTERNAL_USER cannot rewrite key_alias/blocked on someone else's key; existing test covers the owner-can-update-alias case; existing test covers internal-user-cannot-modify-max-budget. --- .../key_management_endpoints.py | 50 ++++++++++---- .../test_key_management_endpoints.py | 68 +++++++++++++++++++ 2 files changed, 106 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index f69d9d2f8d4..32227b3dd6d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1961,22 +1961,48 @@ async def _validate_update_key_data( user_api_key_cache=user_api_key_cache, ) - # Admin-only: only proxy admins, team admins, or org admins can modify max_budget or spend - if ( - data.max_budget is not None and data.max_budget != existing_key_row.max_budget + # Cross-key authorization. Previously only gated on max_budget/spend + # changes, which let a non-admin blanket-rewrite any OTHER field on + # any key (models, alias, metadata, tpm_limit, rpm_limit, + # allowed_routes, guardrails, blocked, duration, permissions, …) as + # long as they avoided budget/spend. + # + # Policy: + # - Key owner (same user_id): may update non-budget fields on their + # own key without the admin check. + # - Anyone else (non-PROXY_ADMIN, not the owner): must pass + # _check_key_admin_access (PROXY_ADMIN / key-owner / team-admin / + # org-admin of the key). + # - max_budget / spend: always require the admin check, even for the + # key owner (matches the existing admin-only budget semantics). + is_key_owner = ( + user_api_key_dict.user_id is not None + and existing_key_row.user_id == user_api_key_dict.user_id + ) + _is_budget_change = ( + data.max_budget is not None + and data.max_budget != existing_key_row.max_budget ) or ( data.spend is not None and data.spend != getattr(existing_key_row, "spend", None) + ) + if ( + (not _is_proxy_admin) + and prisma_client is not None + and (not is_key_owner or _is_budget_change) ): - if prisma_client is not None: - hashed_key = existing_key_row.token - await _check_key_admin_access( - user_api_key_dict=user_api_key_dict, - hashed_token=hashed_key, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - route="/key/update (max_budget/spend)", - ) + hashed_key = existing_key_row.token + await _check_key_admin_access( + user_api_key_dict=user_api_key_dict, + hashed_token=hashed_key, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + route=( + "/key/update (max_budget/spend)" + if _is_budget_change + else "/key/update" + ), + ) # Check team limits if key has a team_id (from request or existing key) team_obj: Optional[LiteLLM_TeamTableCachedObj] = None diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 479defbff5c..41ab60fcf67 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -8030,6 +8030,74 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc assert result is not None +@pytest.mark.asyncio +async def test_update_key_non_budget_rejects_cross_user_modification(monkeypatch): + """Regression: previously _check_key_admin_access was gated on + max_budget/spend changes only, so an internal user could rewrite any + OTHER field (alias, models, tpm_limit, blocked, metadata, …) on any + key they weren't admin of as long as they avoided budget/spend. This + confirms that a non-admin user updating a key that belongs to another + user fails with 403 even for non-budget fields.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + mock_prisma_client = AsyncMock() + test_hashed_token = ( + "cafebabe" * 8 + ) + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = "victim_user" # owned by someone else + mock_existing_key.team_id = None + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "original" + mock_existing_key.models = [] + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": "victim_user", + "team_id": None, + "max_budget": 10.0, + } + + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr( + "litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token + ) + + mock_request = MagicMock() + mock_request.query_params = {} + attacker = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-attacker", + user_id="attacker_user", # NOT the owner + ) + + # Trying to blanket-rewrite a non-budget field on someone else's key + # must now fail. + with pytest.raises(ProxyException) as exc: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest( + key=test_hashed_token, key_alias="pwned", blocked=True + ), + user_api_key_dict=attacker, + litellm_changed_by=None, + ) + assert str(exc.value.code) == "403" + + # ============================================================================ # LIT-1884: Internal users cannot create invalid keys # ============================================================================ From e3b55794ce14e26c8d39b469a5388be4d2f4bc34 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Fri, 17 Apr 2026 00:44:24 +0000 Subject: [PATCH 38/41] fix(proxy): close /organization member_delete + role-escalation gaps Audit-B #7 and #8. 1. /organization/member_delete was not in org_admin_only_routes, so it fell through to management_routes/self_managed_routes and let any caller that reached the route delete arbitrary org memberships without the organization_role_based_access_check that member_add and member_update trigger. Adding it to org_admin_only_routes applies the same ORG_ADMIN-of-target-org gate. 2. /organization/member_update had no validation that the target user was not a global PROXY_ADMIN. An org-admin of any org could alter a PROXY_ADMIN user's per-org role. Reject this unless the caller is PROXY_ADMIN. --- litellm/proxy/_types.py | 7 +++++ .../organization_endpoints.py | 27 +++++++++++++++++++ 2 files changed, 34 insertions(+) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7fff640c497..39857984247 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -690,6 +690,13 @@ class LiteLLMRoutes(enum.Enum): "/organization/delete", "/organization/member_add", "/organization/member_update", + # member_delete is equally destructive as member_add / member_update + # and must be scoped the same way — otherwise it falls through to + # the management_routes / self_managed_routes path and lets any + # non-PROXY_ADMIN caller that reaches the route delete arbitrary + # org memberships without the organization_role_based_access_check + # that member_add / member_update trigger. + "/organization/member_delete", ] # Routes accessible by Admin Viewer (read-only admin access) diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 25df9f0b0f7..670946ac0d8 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -1064,6 +1064,33 @@ async def organization_member_update( }, ) + # Reject attempts to change the role of a global PROXY_ADMIN via + # org-scoped operations. An org-admin of any org could otherwise + # alter a PROXY_ADMIN user's per-org role, which has downstream + # effects on admin UI filtering and scope derivation. + target_user_row = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": data.user_id} + ) + if target_user_row is not None and getattr( + target_user_row, "user_role", None + ) in ( + LitellmUserRoles.PROXY_ADMIN.value, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ): + if ( + user_api_key_dict.user_role + != LitellmUserRoles.PROXY_ADMIN.value + ): + raise HTTPException( + status_code=403, + detail={ + "error": ( + "Only PROXY_ADMIN may modify the organization " + "role of a user who is a global PROXY_ADMIN." + ) + }, + ) + # Update member role if data.role is not None: await prisma_client.db.litellm_organizationmembership.update( From bcb2ea63f139b3ff8bcb2164c3a886c9d5314cea Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 16 Apr 2026 21:12:09 -0700 Subject: [PATCH 39/41] fix: add explicit admin bypass to agent access checks for consistency --- litellm/proxy/agent_endpoints/a2a_routing.py | 23 +++++++++++------- litellm/proxy/agent_endpoints/endpoints.py | 25 ++++++++++++-------- 2 files changed, 29 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index d5f5348731f..a9b75de9807 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -11,7 +11,7 @@ import litellm from fastapi import HTTPException from litellm._logging import verbose_proxy_logger -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth async def route_a2a_agent_request( @@ -50,16 +50,21 @@ async def route_a2a_agent_request( route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) raise ProxyModelNotFoundError(route=route_name, model_name=model_name) - # Verify the caller is permitted to use this agent - is_allowed = await AgentRequestHandler.is_agent_allowed( - agent_id=agent.agent_id, - user_api_key_auth=user_api_key_dict, + # Verify the caller is permitted to use this agent (admins bypass the check) + is_admin = user_api_key_dict is not None and ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value ) - if not is_allowed: - raise HTTPException( - status_code=403, - detail=f"Agent '{agent_name}' is not allowed for your key/team. Contact proxy admin for access.", + if not is_admin: + is_allowed = await AgentRequestHandler.is_agent_allowed( + agent_id=agent.agent_id, + user_api_key_auth=user_api_key_dict, ) + if not is_allowed: + raise HTTPException( + status_code=403, + detail=f"Agent '{agent_name}' is not allowed for your key/team. Contact proxy admin for access.", + ) # Get API base URL from agent config if not agent.agent_card_params or "url" not in agent.agent_card_params: diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 4c5fbaa22c5..c351d3cfecb 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -392,19 +392,24 @@ async def get_agent_by_id( """ await check_feature_access_for_user(user_api_key_dict, "agents") - from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( - AgentRequestHandler, + is_admin = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value ) - - is_allowed = await AgentRequestHandler.is_agent_allowed( - agent_id=agent_id, user_api_key_auth=user_api_key_dict - ) - if not is_allowed: - raise HTTPException( - status_code=403, - detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.", + if not is_admin: + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, ) + is_allowed = await AgentRequestHandler.is_agent_allowed( + agent_id=agent_id, user_api_key_auth=user_api_key_dict + ) + if not is_allowed: + raise HTTPException( + status_code=403, + detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.", + ) + from litellm.proxy.proxy_server import prisma_client if prisma_client is None: From ee2cf0e6e89d3b29aa013292d8dd689822cf306b Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 17 Apr 2026 15:11:45 -0700 Subject: [PATCH 40/41] fix: address three CI failures from recent security PR merges MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - url_utils.py: narrow sockaddr[0] from str|int to str via a helper with a fail-closed isinstance check. Fixes the two mypy errors introduced by the SSRF hardening without masking unexpected stdlib behavior. - key_management_endpoints.py: restore the documented team member_permissions path for /key/update. The cross-key admin check added to close the cross-org rewrite attack was over-broad: it rejected non-admin team members even when can_team_member_execute_key_management_endpoint had already validated their team membership and /key/update grant. Now skip the admin check when the key has a team_id and the change is non-budget (membership + permission already enforced above). Budget/spend changes still require team/org admin. The cross-org attack remains blocked: an outside org admin fails the earlier team membership check. - test_logging_redaction_e2e_test.py: rename and rewrite two parametrized tests to assert that request-body turn_off_message_logging has no effect. Reflects the intentional removal of turn_off_message_logging from _supported_callback_params so the caller cannot override admin logging policy via the request body. - test_key_management_endpoints.py: add two tests covering the restored team member permission path — one positive (non-budget update succeeds for a team member with /key/update grant), one negative (max_budget change still rejected without admin role). --- litellm/litellm_core_utils/url_utils.py | 23 +- .../key_management_endpoints.py | 67 +++--- .../test_logging_redaction_e2e_test.py | 48 ++--- .../test_key_management_endpoints.py | 203 ++++++++++++++++++ 4 files changed, 283 insertions(+), 58 deletions(-) diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index b55882819de..a65d0892aa2 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -78,6 +78,22 @@ def _format_host_header(hostname: str, port: int, default_port: int) -> str: return f"{bracketed}:{port}" +def _sockaddr_host(sockaddr: Any) -> str: + """Return the host element of a ``getaddrinfo`` sockaddr as ``str``. + + ``getaddrinfo`` with ``IPPROTO_TCP`` returns AF_INET / AF_INET6 sockaddrs + whose first element is always a host string. mypy types it as + ``str | int`` (since sockaddrs for other families can hold ints), so we + narrow at the boundary. Fail closed if the stdlib ever returns something + unexpected — a non-string here would mean we have no IP to check against + the SSRF blocklist. + """ + host = sockaddr[0] + if not isinstance(host, str): + raise SSRFError(f"getaddrinfo returned non-string host: {host!r}") + return host + + def _is_host_allowlisted(hostname: str, effective_port: int) -> bool: """Check whether a host is in the admin-configured allowlist. @@ -148,9 +164,10 @@ def validate_url(url: str) -> Tuple[str, str]: if not is_allowlisted: for family, type_, proto, canonname, sockaddr in addrinfo: - if _is_blocked_ip(sockaddr[0]): + resolved_ip = _sockaddr_host(sockaddr) + if _is_blocked_ip(resolved_ip): raise SSRFError( - f"URL targets a blocked address ({sockaddr[0]}). " + f"URL targets a blocked address ({resolved_ip}). " "If this is a legitimate internal service, add the host " "to `user_url_allowed_hosts` in general_settings." ) @@ -166,7 +183,7 @@ def validate_url(url: str) -> Tuple[str, str]: # For HTTP, rewrite URL to connect to the validated IP directly # to prevent DNS rebinding (no TLS to bind the connection). - validated_ip = addrinfo[0][4][0] + validated_ip = _sockaddr_host(addrinfo[0][4]) is_ipv6 = addrinfo[0][0] == socket.AF_INET6 ip_host = f"[{validated_ip}]" if is_ipv6 else validated_ip diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 32227b3dd6d..98e41c409b5 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -768,9 +768,9 @@ async def _common_key_generation_helper( # noqa: PLR0915 request_type="key", **data_json, table_name="key" ) - response[ - "soft_budget" - ] = data.soft_budget # include the user-input soft budget in the response + response["soft_budget"] = ( + data.soft_budget + ) # include the user-input soft budget in the response response = GenerateKeyResponse(**response) @@ -1970,26 +1970,37 @@ async def _validate_update_key_data( # Policy: # - Key owner (same user_id): may update non-budget fields on their # own key without the admin check. - # - Anyone else (non-PROXY_ADMIN, not the owner): must pass - # _check_key_admin_access (PROXY_ADMIN / key-owner / team-admin / - # org-admin of the key). + # - Team member with /key/update grant (on a team key): may update + # non-budget fields. Team membership + permission is already + # enforced by can_team_member_execute_key_management_endpoint + # above, which raises 401 for non-members or members without the + # grant — so reaching this point on a team key means the caller + # was authorized via member_permissions. This preserves the + # documented member_permissions feature while still blocking the + # cross-org attack (an outside org admin is not a member of the + # victim team and gets rejected at the earlier check). + # - Anyone else (non-PROXY_ADMIN, not the owner, not a team member + # on a team key): must pass _check_key_admin_access (PROXY_ADMIN + # / key-owner / team-admin / org-admin of the key). # - max_budget / spend: always require the admin check, even for the - # key owner (matches the existing admin-only budget semantics). + # key owner or a team member (matches the existing admin-only + # budget semantics). is_key_owner = ( user_api_key_dict.user_id is not None and existing_key_row.user_id == user_api_key_dict.user_id ) _is_budget_change = ( - data.max_budget is not None - and data.max_budget != existing_key_row.max_budget + data.max_budget is not None and data.max_budget != existing_key_row.max_budget ) or ( data.spend is not None and data.spend != getattr(existing_key_row, "spend", None) ) + is_team_key = existing_key_row.team_id is not None + can_skip_admin_check_for_non_budget = is_key_owner or is_team_key if ( (not _is_proxy_admin) and prisma_client is not None - and (not is_key_owner or _is_budget_change) + and (_is_budget_change or not can_skip_admin_check_for_non_budget) ): hashed_key = existing_key_row.token await _check_key_admin_access( @@ -1998,9 +2009,7 @@ async def _validate_update_key_data( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, route=( - "/key/update (max_budget/spend)" - if _is_budget_change - else "/key/update" + "/key/update (max_budget/spend)" if _is_budget_change else "/key/update" ), ) @@ -3303,10 +3312,10 @@ async def delete_verification_tokens( try: if prisma_client: tokens = [_hash_token_if_needed(token=key) for key in tokens] - _keys_being_deleted: List[ - LiteLLM_VerificationToken - ] = await prisma_client.db.litellm_verificationtoken.find_many( - where={"token": {"in": tokens}} + _keys_being_deleted: List[LiteLLM_VerificationToken] = ( + await prisma_client.db.litellm_verificationtoken.find_many( + where={"token": {"in": tokens}} + ) ) if len(_keys_being_deleted) == 0: @@ -3506,9 +3515,9 @@ async def _rotate_master_key( # noqa: PLR0915 from litellm.proxy.proxy_server import proxy_config try: - models: Optional[ - List - ] = await prisma_client.db.litellm_proxymodeltable.find_many() + models: Optional[List] = ( + await prisma_client.db.litellm_proxymodeltable.find_many() + ) except Exception: models = None # 2. process model table @@ -4148,11 +4157,11 @@ async def validate_key_list_check( param="user_id", code=status.HTTP_403_FORBIDDEN, ) - complete_user_info_db_obj: Optional[ - BaseModel - ] = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_api_key_dict.user_id}, - include={"organization_memberships": True}, + complete_user_info_db_obj: Optional[BaseModel] = ( + await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user_api_key_dict.user_id}, + include={"organization_memberships": True}, + ) ) if complete_user_info_db_obj is None: @@ -4235,10 +4244,10 @@ async def _fetch_user_team_objects( if complete_user_info is None or not complete_user_info.teams: return [] - teams: Optional[ - List[BaseModel] - ] = await prisma_client.db.litellm_teamtable.find_many( - where={"team_id": {"in": complete_user_info.teams}} + teams: Optional[List[BaseModel]] = ( + await prisma_client.db.litellm_teamtable.find_many( + where={"team_id": {"in": complete_user_info.teams}} + ) ) if teams is None: return [] diff --git a/tests/logging_callback_tests/test_logging_redaction_e2e_test.py b/tests/logging_callback_tests/test_logging_redaction_e2e_test.py index 0391a5a8957..3047578dbb3 100644 --- a/tests/logging_callback_tests/test_logging_redaction_e2e_test.py +++ b/tests/logging_callback_tests/test_logging_redaction_e2e_test.py @@ -56,7 +56,13 @@ async def test_global_redaction_on(): @pytest.mark.parametrize("turn_off_message_logging", [True, False]) @pytest.mark.asyncio -async def test_global_redaction_with_dynamic_params(turn_off_message_logging): +async def test_global_redaction_ignores_dynamic_param(turn_off_message_logging): + """ + Request-body `turn_off_message_logging` is no longer honored as a dynamic + callback param — global setting (or admin-configured key/team config) wins. + With global redaction ON, the caller cannot disable redaction via the + request body. + """ litellm.turn_off_message_logging = True test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] @@ -75,23 +81,20 @@ async def test_global_redaction_with_dynamic_params(turn_off_message_logging): json.dumps(standard_logging_payload, indent=2), ) - if turn_off_message_logging is True: - response = standard_logging_payload["response"] - assert response["choices"][0]["message"]["content"] == "redacted-by-litellm" - assert ( - standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm" - ) - else: - assert ( - standard_logging_payload["response"]["choices"][0]["message"]["content"] - == "hello" - ) - assert standard_logging_payload["messages"][0]["content"] == "hi" + response = standard_logging_payload["response"] + assert response["choices"][0]["message"]["content"] == "redacted-by-litellm" + assert standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm" @pytest.mark.parametrize("turn_off_message_logging", [True, False]) @pytest.mark.asyncio -async def test_global_redaction_off_with_dynamic_params(turn_off_message_logging): +async def test_global_redaction_off_ignores_dynamic_param(turn_off_message_logging): + """ + Request-body `turn_off_message_logging` is no longer honored as a dynamic + callback param — global setting (or admin-configured key/team config) wins. + With global redaction OFF, the caller cannot enable redaction via the + request body. + """ litellm.turn_off_message_logging = False test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] @@ -109,18 +112,11 @@ async def test_global_redaction_off_with_dynamic_params(turn_off_message_logging "logged standard logging payload", json.dumps(standard_logging_payload, indent=2), ) - if turn_off_message_logging is True: - response = standard_logging_payload["response"] - assert response["choices"][0]["message"]["content"] == "redacted-by-litellm" - assert ( - standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm" - ) - else: - assert ( - standard_logging_payload["response"]["choices"][0]["message"]["content"] - == "hello" - ) - assert standard_logging_payload["messages"][0]["content"] == "hi" + assert ( + standard_logging_payload["response"]["choices"][0]["message"]["content"] + == "hello" + ) + assert standard_logging_payload["messages"][0]["content"] == "hi" @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 41ab60fcf67..811fd8d5cfa 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -8098,6 +8098,209 @@ async def test_update_key_non_budget_rejects_cross_user_modification(monkeypatch assert str(exc.value.code) == "403" +@pytest.mark.asyncio +async def test_update_key_team_member_with_permission_can_update_non_budget( + monkeypatch, +): + """A team member whose team grants /key/update in member_permissions can + update non-budget fields on a team key even though they are not a team + admin. Regression: the cross-key admin check was over-broad and rejected + this documented path.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + test_hashed_token = "deadbeef" * 8 + team_id = "team-with-update-grant" + member_user_id = "team-member-user" + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = None # team-scoped key (no owning user) + mock_existing_key.team_id = team_id + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "original" + mock_existing_key.models = [] + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": None, + "team_id": team_id, + "max_budget": 10.0, + } + + team_table = LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="some-team-admin", role="admin"), + Member(user_id=member_user_id, role="user"), + ], + team_member_permissions=["/key/update", "/key/info"], + ) + + mock_updated_key = MagicMock() + mock_updated_key.token = test_hashed_token + mock_updated_key.key_alias = "renamed-by-member" + + mock_prisma_client = AsyncMock() + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.update_data = AsyncMock(return_value=mock_updated_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + async def mock_get_team_object(*args, **kwargs): + return team_table + + async def mock_enforce_unique_key_alias(**kwargs): + pass + + async def mock_delete_cache_key_object(**kwargs): + pass + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias", + mock_enforce_unique_key_alias, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + mock_delete_cache_key_object, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + monkeypatch.setattr( + "litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token + ) + + mock_request = MagicMock() + mock_request.query_params = {} + team_member = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-member", + user_id=member_user_id, + team_id=team_id, + ) + + # Non-budget update on a team key by a team member with /key/update + # permission should succeed. + result = await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, key_alias="renamed-by-member"), + user_api_key_dict=team_member, + litellm_changed_by=None, + ) + + assert result is not None + + +@pytest.mark.asyncio +async def test_update_key_team_member_cannot_change_budget(monkeypatch): + """A team member with /key/update in member_permissions still cannot + change max_budget — budget/spend changes require team/org admin. The + member_permissions bypass only applies to non-budget fields.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + test_hashed_token = "feedface" * 8 + team_id = "team-with-update-grant" + member_user_id = "team-member-user" + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = None # team-scoped key (no owning user) + mock_existing_key.team_id = team_id + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "original" + mock_existing_key.models = [] + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": None, + "team_id": team_id, + "max_budget": 10.0, + } + + team_table = LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="some-team-admin", role="admin"), + Member(user_id=member_user_id, role="user"), + ], + team_member_permissions=["/key/update", "/key/info"], + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr( + "litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token + ) + + mock_request = MagicMock() + mock_request.query_params = {} + team_member = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-member", + user_id=member_user_id, + team_id=team_id, + ) + + with pytest.raises(ProxyException) as exc: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, max_budget=500.0), + user_api_key_dict=team_member, + litellm_changed_by=None, + ) + assert str(exc.value.code) == "403" + + # ============================================================================ # LIT-1884: Internal users cannot create invalid keys # ============================================================================ From 9c0b73e5f41f4f3e598d8fdc38c64910d80ce295 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 18 Apr 2026 11:00:09 -0700 Subject: [PATCH 41/41] [Fix] should_create_missing_views returns False for reltuples=0 (falsy zero bug) `should_create_missing_views()` had `and result[0]["reltuples"]` which is falsy when reltuples=0. On a fresh empty PostgreSQL table, CREATE INDEX sets reltuples=0, causing the guard to return False and skip view creation entirely. Views like MonthlyGlobalSpendPerKey are never created, and the /global/spend/logs endpoint returns 500. Fix: change to `and result[0]["reltuples"] is not None` so reltuples=0 (empty table) and reltuples=-1 (unanalyzed table) both correctly return True. Also harden test_vertex_ai.py to return None instead of crashing with JSONDecodeError when the spend-logs endpoint returns a non-JSON 500 response, and add unit tests covering all three reltuples branches (0, -1, positive). --- litellm/proxy/db/create_views.py | 2 +- tests/pass_through_tests/test_vertex_ai.py | 4 +++ .../proxy/db/test_create_views.py | 36 +++++++++++++++++++ 3 files changed, 41 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/db/create_views.py b/litellm/proxy/db/create_views.py index 3598045545b..d84cebcf05a 100644 --- a/litellm/proxy/db/create_views.py +++ b/litellm/proxy/db/create_views.py @@ -251,7 +251,7 @@ async def should_create_missing_views(db: _db) -> bool: and len(result) > 0 and isinstance(result[0], dict) and "reltuples" in result[0] - and result[0]["reltuples"] + and result[0]["reltuples"] is not None and (result[0]["reltuples"] == 0 or result[0]["reltuples"] == -1) ): verbose_logger.debug("Should create views") diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index ba27a4cc460..73bf03c5000 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -72,6 +72,10 @@ async def call_spend_logs_endpoint(): response = requests.get(url, headers=headers) print("response from call_spend_logs_endpoint", response) + if response.status_code != 200: + print(f"spend logs endpoint returned {response.status_code}: {response.text}") + return None + json_response = response.json() # get spend for today diff --git a/tests/test_litellm/proxy/db/test_create_views.py b/tests/test_litellm/proxy/db/test_create_views.py index 1a90b4c204d..c0c09d0137b 100644 --- a/tests/test_litellm/proxy/db/test_create_views.py +++ b/tests/test_litellm/proxy/db/test_create_views.py @@ -130,6 +130,42 @@ async def test_create_views_reraises_undefined_function_error(): mock_db.execute_raw.assert_not_called() +@pytest.mark.asyncio +async def test_should_create_missing_views_reltuples_zero(): + """should return True when reltuples is 0 (fresh empty table).""" + from litellm.proxy.db.create_views import should_create_missing_views + + mock_db = MagicMock() + mock_db.query_raw = AsyncMock(return_value=[{"reltuples": 0}]) + + result = await should_create_missing_views(mock_db) + assert result is True + + +@pytest.mark.asyncio +async def test_should_create_missing_views_reltuples_negative_one(): + """should return True when reltuples is -1 (table created, no ANALYZE yet).""" + from litellm.proxy.db.create_views import should_create_missing_views + + mock_db = MagicMock() + mock_db.query_raw = AsyncMock(return_value=[{"reltuples": -1}]) + + result = await should_create_missing_views(mock_db) + assert result is True + + +@pytest.mark.asyncio +async def test_should_create_missing_views_reltuples_positive(): + """should return False when reltuples > 0 (table has data).""" + from litellm.proxy.db.create_views import should_create_missing_views + + mock_db = MagicMock() + mock_db.query_raw = AsyncMock(return_value=[{"reltuples": 1000}]) + + result = await should_create_missing_views(mock_db) + assert result is False + + @pytest.mark.asyncio async def test_create_views_creates_view_on_undefined_table_error(): """should treat 'undefined table' as a missing-view signal and attempt creation."""