From e3d1e0345cdaa258ee220c67ea322dd3896181c3 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 13 Jan 2026 17:01:51 -0800 Subject: [PATCH 1/9] only show own internal user usage --- .../common_daily_activity.py | 9 +- .../management_endpoints/team_endpoints.py | 34 +- .../test_team_endpoints.py | 363 ++++++++++++++++++ 3 files changed, 401 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index f52abf86b97..c52491efc7c 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -343,7 +343,7 @@ def _build_where_conditions( start_date: str, end_date: str, model: Optional[str], - api_key: Optional[str], + api_key: Optional[Union[str, List[str]]], exclude_entity_ids: Optional[List[str]] = None, ) -> Dict[str, Any]: """Build prisma where clause for daily activity queries.""" @@ -357,7 +357,10 @@ def _build_where_conditions( if model: where_conditions["model"] = model if api_key: - where_conditions["api_key"] = api_key + if isinstance(api_key, list): + where_conditions["api_key"] = {"in": api_key} + else: + where_conditions["api_key"] = api_key if entity_id is not None: if isinstance(entity_id, list): @@ -445,7 +448,7 @@ async def get_daily_activity( start_date: Optional[str], end_date: Optional[str], model: Optional[str], - api_key: Optional[str], + api_key: Optional[Union[str, List[str]]], page: int, page_size: int, exclude_entity_ids: Optional[List[str]] = None, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 78caa86db7b..d1549b51167 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -3601,7 +3601,7 @@ async def get_team_daily_activity( }, ) - ## Fetch team aliases + ## Fetch team aliases and check team admin status where_condition = {} if team_ids_list: where_condition["team_id"] = {"in": list(team_ids_list)} @@ -3612,6 +3612,36 @@ async def get_team_daily_activity( t.team_id: {"team_alias": t.team_alias} for t in team_aliases } + # Check if user is team admin for any requested teams + # If not, filter by user's API keys + user_api_keys: Optional[List[str]] = None + if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: + # Check if user is team admin for any of the teams + is_team_admin_for_any = False + for team_alias in team_aliases: + team_obj = LiteLLM_TeamTable(**team_alias.model_dump()) + if _is_user_team_admin( + user_api_key_dict=user_api_key_dict, team_obj=team_obj + ): + is_team_admin_for_any = True + break + + # If user is not a team admin for any team, filter by their API keys + if not is_team_admin_for_any: + # Get all API keys for this user + user_keys = await prisma_client.db.litellm_verificationtoken.find_many( + where={"user_id": user_api_key_dict.user_id} + ) + user_api_keys = [key.token for key in user_keys if key.token] + # If user has no API keys, return empty result + if not user_api_keys: + user_api_keys = [""] # Use empty string to ensure no matches + + # If api_key parameter is provided, use it; otherwise use user_api_keys if set + final_api_key_filter: Optional[Union[str, List[str]]] = api_key + if final_api_key_filter is None and user_api_keys is not None: + final_api_key_filter = user_api_keys + return await get_daily_activity( prisma_client=prisma_client, table_name="litellm_dailyteamspend", @@ -3622,7 +3652,7 @@ async def get_team_daily_activity( start_date=start_date, end_date=end_date, model=model, - api_key=api_key, + api_key=final_api_key_filter, page=page, page_size=page_size, ) 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 e296066b998..bbff7448e13 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -20,6 +20,7 @@ from litellm.proxy._types import ( LiteLLM_OrganizationTable, LiteLLM_OrganizationTableWithMembers, LiteLLM_TeamTable, + LiteLLM_UserTable, LitellmUserRoles, Member, ProxyErrorTypes, @@ -4476,6 +4477,187 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth): assert deserialized_settings == router_settings_data +@pytest.mark.asyncio +async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys( + mock_db_client, +): + """ + Test that non-team-admin users only see their own spend (filtered by their API keys) + when calling /team/daily/activity endpoint. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + get_team_daily_activity, + ) + + # Create a non-admin user + user_id = "test_user_123" + team_id = "test_team_456" + user_api_key_dict = UserAPIKeyAuth( + user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER + ) + + # Mock user info + mock_user_info = LiteLLM_UserTable( + user_id=user_id, + teams=[team_id], + max_budget=1000.0, + spend=0.0, + user_email="test@example.com", + user_role="internal_user", + ) + + # Mock team with user as non-admin member + mock_team_member = Member(user_id=user_id, role="user") + mock_team = MagicMock(spec=LiteLLM_TeamTable) + mock_team.team_id = team_id + mock_team.team_alias = "Test Team" + mock_team.members_with_roles = [mock_team_member] + mock_team.model_dump.return_value = { + "team_id": team_id, + "team_alias": "Test Team", + "members_with_roles": [{"user_id": user_id, "role": "user"}], + } + + # Mock user's API keys + user_api_key_1 = MagicMock() + user_api_key_1.token = "user_key_1" + user_api_key_2 = MagicMock() + user_api_key_2.token = "user_key_2" + + # Setup mocks + mock_db_client.db.litellm_teamtable.find_many = AsyncMock( + return_value=[mock_team] + ) + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[user_api_key_1, user_api_key_2] + ) + + # Mock get_user_object + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + ) as mock_get_user_object: + mock_get_user_object.return_value = mock_user_info + + # Mock get_daily_activity to capture the api_key parameter + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity: + mock_get_daily_activity.return_value = MagicMock() + + # Call the endpoint + await get_team_daily_activity( + team_ids=team_id, + start_date="2024-01-01", + end_date="2024-01-02", + model=None, + api_key=None, + page=1, + page_size=10, + exclude_team_ids=None, + user_api_key_dict=user_api_key_dict, + ) + + # Verify get_daily_activity was called with user's API keys as filter + mock_get_daily_activity.assert_called_once() + call_kwargs = mock_get_daily_activity.call_args[1] + assert call_kwargs["api_key"] == ["user_key_1", "user_key_2"] + assert call_kwargs["entity_id"] == [team_id] + + # Verify user's API keys were fetched + mock_db_client.db.litellm_verificationtoken.find_many.assert_called_once() + api_key_call_kwargs = ( + mock_db_client.db.litellm_verificationtoken.find_many.call_args[1] + ) + assert api_key_call_kwargs["where"] == {"user_id": user_id} + + +@pytest.mark.asyncio +async def test_get_team_daily_activity_team_admin_sees_all_spend(mock_db_client): + """ + Test that team admin users see all team spend (no API key filtering) + when calling /team/daily/activity endpoint. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + get_team_daily_activity, + ) + + # Create a team admin user + user_id = "test_admin_123" + team_id = "test_team_456" + user_api_key_dict = UserAPIKeyAuth( + user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER + ) + + # Mock user info + mock_user_info = LiteLLM_UserTable( + user_id=user_id, + teams=[team_id], + max_budget=1000.0, + spend=0.0, + user_email="admin@example.com", + user_role="internal_user", + ) + + # Mock team with user as admin member + mock_team_member = Member(user_id=user_id, role="admin") + mock_team = MagicMock(spec=LiteLLM_TeamTable) + mock_team.team_id = team_id + mock_team.team_alias = "Test Team" + mock_team.members_with_roles = [mock_team_member] + mock_team.model_dump.return_value = { + "team_id": team_id, + "team_alias": "Test Team", + "members_with_roles": [{"user_id": user_id, "role": "admin"}], + } + + # Setup mocks + mock_db_client.db.litellm_teamtable.find_many = AsyncMock( + return_value=[mock_team] + ) + + # Mock get_user_object + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + ) as mock_get_user_object: + mock_get_user_object.return_value = mock_user_info + + # Mock get_daily_activity to capture the api_key parameter + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity: + mock_get_daily_activity.return_value = MagicMock() + + # Call the endpoint + await get_team_daily_activity( + team_ids=team_id, + start_date="2024-01-01", + end_date="2024-01-02", + model=None, + api_key=None, + page=1, + page_size=10, + exclude_team_ids=None, + user_api_key_dict=user_api_key_dict, + ) + + # Verify get_daily_activity was called WITHOUT API key filtering + mock_get_daily_activity.assert_called_once() + call_kwargs = mock_get_daily_activity.call_args[1] + assert call_kwargs["api_key"] is None + assert call_kwargs["entity_id"] == [team_id] + + # Verify user's API keys were NOT fetched (since they're admin) + if hasattr( + mock_db_client.db.litellm_verificationtoken, "find_many" + ) and mock_db_client.db.litellm_verificationtoken.find_many.called: + # If it was called, that's unexpected for admin users + assert False, "API keys should not be fetched for team admin users" + + @pytest.mark.asyncio async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth): """ @@ -4552,3 +4734,184 @@ async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth) # Verify router_settings can be deserialized and matches input deserialized_settings = json.loads(team_data["router_settings"]) assert deserialized_settings == router_settings_data + + +@pytest.mark.asyncio +async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys( + mock_db_client, +): + """ + Test that non-team-admin users only see their own spend (filtered by their API keys) + when calling /team/daily/activity endpoint. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + get_team_daily_activity, + ) + + # Create a non-admin user + user_id = "test_user_123" + team_id = "test_team_456" + user_api_key_dict = UserAPIKeyAuth( + user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER + ) + + # Mock user info + mock_user_info = LiteLLM_UserTable( + user_id=user_id, + teams=[team_id], + max_budget=1000.0, + spend=0.0, + user_email="test@example.com", + user_role="internal_user", + ) + + # Mock team with user as non-admin member + mock_team_member = Member(user_id=user_id, role="user") + mock_team = MagicMock(spec=LiteLLM_TeamTable) + mock_team.team_id = team_id + mock_team.team_alias = "Test Team" + mock_team.members_with_roles = [mock_team_member] + mock_team.model_dump.return_value = { + "team_id": team_id, + "team_alias": "Test Team", + "members_with_roles": [{"user_id": user_id, "role": "user"}], + } + + # Mock user's API keys + user_api_key_1 = MagicMock() + user_api_key_1.token = "user_key_1" + user_api_key_2 = MagicMock() + user_api_key_2.token = "user_key_2" + + # Setup mocks + mock_db_client.db.litellm_teamtable.find_many = AsyncMock( + return_value=[mock_team] + ) + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[user_api_key_1, user_api_key_2] + ) + + # Mock get_user_object + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + ) as mock_get_user_object: + mock_get_user_object.return_value = mock_user_info + + # Mock get_daily_activity to capture the api_key parameter + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity: + mock_get_daily_activity.return_value = MagicMock() + + # Call the endpoint + await get_team_daily_activity( + team_ids=team_id, + start_date="2024-01-01", + end_date="2024-01-02", + model=None, + api_key=None, + page=1, + page_size=10, + exclude_team_ids=None, + user_api_key_dict=user_api_key_dict, + ) + + # Verify get_daily_activity was called with user's API keys as filter + mock_get_daily_activity.assert_called_once() + call_kwargs = mock_get_daily_activity.call_args[1] + assert call_kwargs["api_key"] == ["user_key_1", "user_key_2"] + assert call_kwargs["entity_id"] == [team_id] + + # Verify user's API keys were fetched + mock_db_client.db.litellm_verificationtoken.find_many.assert_called_once() + api_key_call_kwargs = ( + mock_db_client.db.litellm_verificationtoken.find_many.call_args[1] + ) + assert api_key_call_kwargs["where"] == {"user_id": user_id} + + +@pytest.mark.asyncio +async def test_get_team_daily_activity_team_admin_sees_all_spend(mock_db_client): + """ + Test that team admin users see all team spend (no API key filtering) + when calling /team/daily/activity endpoint. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + get_team_daily_activity, + ) + + # Create a team admin user + user_id = "test_admin_123" + team_id = "test_team_456" + user_api_key_dict = UserAPIKeyAuth( + user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER + ) + + # Mock user info + mock_user_info = LiteLLM_UserTable( + user_id=user_id, + teams=[team_id], + max_budget=1000.0, + spend=0.0, + user_email="admin@example.com", + user_role="internal_user", + ) + + # Mock team with user as admin member + mock_team_member = Member(user_id=user_id, role="admin") + mock_team = MagicMock(spec=LiteLLM_TeamTable) + mock_team.team_id = team_id + mock_team.team_alias = "Test Team" + mock_team.members_with_roles = [mock_team_member] + mock_team.model_dump.return_value = { + "team_id": team_id, + "team_alias": "Test Team", + "members_with_roles": [{"user_id": user_id, "role": "admin"}], + } + + # Setup mocks + mock_db_client.db.litellm_teamtable.find_many = AsyncMock( + return_value=[mock_team] + ) + + # Mock get_user_object + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new_callable=AsyncMock, + ) as mock_get_user_object: + mock_get_user_object.return_value = mock_user_info + + # Mock get_daily_activity to capture the api_key parameter + with patch( + "litellm.proxy.management_endpoints.team_endpoints.get_daily_activity", + new_callable=AsyncMock, + ) as mock_get_daily_activity: + mock_get_daily_activity.return_value = MagicMock() + + # Call the endpoint + await get_team_daily_activity( + team_ids=team_id, + start_date="2024-01-01", + end_date="2024-01-02", + model=None, + api_key=None, + page=1, + page_size=10, + exclude_team_ids=None, + user_api_key_dict=user_api_key_dict, + ) + + # Verify get_daily_activity was called WITHOUT API key filtering + mock_get_daily_activity.assert_called_once() + call_kwargs = mock_get_daily_activity.call_args[1] + assert call_kwargs["api_key"] is None + assert call_kwargs["entity_id"] == [team_id] + + # Verify user's API keys were NOT fetched (since they're admin) + if hasattr( + mock_db_client.db.litellm_verificationtoken, "find_many" + ) and mock_db_client.db.litellm_verificationtoken.find_many.called: + # If it was called, that's unexpected for admin users + assert False, "API keys should not be fetched for team admin users" From 27a246722630653cb46f45ceee06d5ee44286ef3 Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Sat, 17 Jan 2026 00:15:35 +0530 Subject: [PATCH 2/9] fix: correct budget limit validation operator (>=) for team members (#19207) --- litellm/proxy/auth/auth_checks.py | 138 ++++++++++++++++-------------- 1 file changed, 73 insertions(+), 65 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index a741869e5fc..5e0a211906e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -202,21 +202,29 @@ async def common_checks( and general_settings["enforce_user_param"] is True ): # Get HTTP method from request - http_method = request.method if hasattr(request, 'method') else None - + http_method = request.method if hasattr(request, "method") else None + # Check if it's a POST request and if it's an OpenAI route but not MCP is_post_method = http_method and http_method.upper() == "POST" is_openai_route = RouteChecks.is_llm_api_route(route=route) - is_mcp_route = route in LiteLLMRoutes.mcp_routes.value or RouteChecks.check_route_access( - route=route, allowed_routes=LiteLLMRoutes.mcp_routes.value + is_mcp_route = ( + route in LiteLLMRoutes.mcp_routes.value + or RouteChecks.check_route_access( + route=route, allowed_routes=LiteLLMRoutes.mcp_routes.value + ) ) - + # Enforce user param only for POST requests on OpenAI routes (excluding MCP routes) - if is_post_method and is_openai_route and not is_mcp_route and "user" not in request_body: + if ( + is_post_method + and is_openai_route + and not is_mcp_route + and "user" not in request_body + ): raise Exception( f"'user' param not passed in. 'enforce_user_param'={general_settings['enforce_user_param']}" ) - + # 6.1 [OPTIONAL] If 'reject_clientside_metadata_tags' enabled - reject request if it has client-side 'metadata.tags' if ( general_settings.get("reject_clientside_metadata_tags", None) is not None @@ -502,53 +510,51 @@ async def get_default_end_user_budget( ) -> Optional[LiteLLM_BudgetTable]: """ Fetches the default end user budget from the database if litellm.max_end_user_budget_id is configured. - + This budget is applied to end users who don't have an explicit budget_id set. Results are cached for performance. - + Args: prisma_client: Database client instance user_api_key_cache: Cache for storing/retrieving budget data parent_otel_span: Optional OpenTelemetry span for tracing - + Returns: LiteLLM_BudgetTable if configured and found, None otherwise """ if prisma_client is None or litellm.max_end_user_budget_id is None: return None - + cache_key = f"default_end_user_budget:{litellm.max_end_user_budget_id}" - + # Check cache first cached_budget = await user_api_key_cache.async_get_cache(key=cache_key) if cached_budget is not None: return LiteLLM_BudgetTable(**cached_budget) - + # Fetch from database try: budget_record = await prisma_client.db.litellm_budgettable.find_unique( where={"budget_id": litellm.max_end_user_budget_id} ) - + if budget_record is None: verbose_proxy_logger.warning( f"Default end user budget not found in database: {litellm.max_end_user_budget_id}" ) return None - + # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( - key=cache_key, + key=cache_key, value=budget_record.dict(), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) - + return LiteLLM_BudgetTable(**budget_record.dict()) - + except Exception as e: - verbose_proxy_logger.error( - f"Error fetching default end user budget: {str(e)}" - ) + verbose_proxy_logger.error(f"Error fetching default end user budget: {str(e)}") return None @@ -560,38 +566,38 @@ async def _apply_default_budget_to_end_user( ) -> LiteLLM_EndUserTable: """ Helper function to apply default budget to end user if they don't have a budget assigned. - + Args: end_user_obj: The end user object to potentially apply default budget to prisma_client: Database client instance user_api_key_cache: Cache for storing/retrieving data parent_otel_span: Optional OpenTelemetry span for tracing - + Returns: Updated end user object with default budget applied if applicable """ # If end user already has a budget assigned, no need to apply default if end_user_obj.litellm_budget_table is not None: return end_user_obj - + # If no default budget configured, return as-is if litellm.max_end_user_budget_id is None: return end_user_obj - + # Fetch and apply default budget default_budget = await get_default_end_user_budget( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, ) - + if default_budget is not None: # Apply default budget to end user object end_user_obj.litellm_budget_table = default_budget verbose_proxy_logger.debug( f"Applied default budget {litellm.max_end_user_budget_id} to end user {end_user_obj.user_id}" ) - + return end_user_obj @@ -601,20 +607,20 @@ def _check_end_user_budget( ) -> None: """ Check if end user is within their budget limit. - + Args: end_user_obj: The end user object to check route: The request route - + Raises: litellm.BudgetExceededError: If end user has exceeded their budget """ if route in LiteLLMRoutes.info_routes.value: return - + if end_user_obj.litellm_budget_table is None: return - + end_user_budget = end_user_obj.litellm_budget_table.max_budget if end_user_budget is not None and end_user_obj.spend > end_user_budget: raise litellm.BudgetExceededError( @@ -635,8 +641,8 @@ async def get_end_user_object( ) -> Optional[LiteLLM_EndUserTable]: """ Returns end user object from database or cache. - - If end user exists but has no budget_id, applies the default budget + + If end user exists but has no budget_id, applies the default budget (if configured via litellm.max_end_user_budget_id). Args: @@ -646,7 +652,7 @@ async def get_end_user_object( route: The request route parent_otel_span: Optional OpenTelemetry span for tracing proxy_logging_obj: Optional proxy logging object - + Returns: LiteLLM_EndUserTable if found, None otherwise """ @@ -655,14 +661,14 @@ async def get_end_user_object( if end_user_id is None: return None - + _key = "end_user_id:{}".format(end_user_id) # Check cache first cached_user_obj = await user_api_key_cache.async_get_cache(key=_key) if cached_user_obj is not None: return_obj = LiteLLM_EndUserTable(**cached_user_obj) - + # Apply default budget if needed return_obj = await _apply_default_budget_to_end_user( end_user_obj=return_obj, @@ -670,10 +676,10 @@ async def get_end_user_object( user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, ) - + # Check budget limits _check_end_user_budget(end_user_obj=return_obj, route=route) - + return return_obj # Fetch from database @@ -688,7 +694,7 @@ async def get_end_user_object( # Convert to LiteLLM_EndUserTable object _response = LiteLLM_EndUserTable(**response.dict()) - + # Apply default budget if needed _response = await _apply_default_budget_to_end_user( end_user_obj=_response, @@ -696,18 +702,17 @@ async def get_end_user_object( user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, ) - + # Save to cache (always store as dict for consistency) await user_api_key_cache.async_set_cache( - key="end_user_id:{}".format(end_user_id), - value=_response.dict() + key="end_user_id:{}".format(end_user_id), value=_response.dict() ) - + # Check budget limits _check_end_user_budget(end_user_obj=_response, route=route) return _response - + except Exception as e: if isinstance(e, litellm.BudgetExceededError): raise e @@ -747,7 +752,6 @@ async def get_tag_objects_batch( tag_objects = {} uncached_tags = [] - # Try to get all tags from cache first for tag_name in tag_names: @@ -1138,7 +1142,6 @@ async def _cache_management_object( user_api_key_cache: DualCache, proxy_logging_obj: Optional[ProxyLogging], ): - await user_api_key_cache.async_set_cache( key=key, value=value, @@ -1459,9 +1462,7 @@ async def get_team_object_by_alias( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception( - "Error looking up team by alias: %s", team_alias - ) + verbose_proxy_logger.exception("Error looking up team by alias: %s", team_alias) raise HTTPException( status_code=500, detail={ @@ -1602,11 +1603,11 @@ class ExperimentalUIJWTToken: ) -> str: """ Generate a JWT token for CLI authentication with 24-hour expiration. - + Args: user_info: User information from the database team_id: Team ID for the user (optional, uses user's team if available) - + Returns: Encrypted JWT token string """ @@ -1800,7 +1801,7 @@ async def get_org_object( - Check if org id in proxy Org Table - if valid, return LiteLLM_OrganizationTable object - if not, then raise an error - + Args: org_id: Organization ID to look up prisma_client: Database client @@ -1820,7 +1821,7 @@ async def get_org_object( cache_key = "org_id:{}".format(org_id) if include_budget_table: cache_key = "org_id:{}:with_budget".format(org_id) - + # check if in cache cached_org_obj = user_api_key_cache.async_get_cache(key=cache_key) if cached_org_obj is not None: @@ -1833,7 +1834,7 @@ async def get_org_object( query_kwargs: Dict[str, Any] = {"where": {"organization_id": org_id}} if include_budget_table: query_kwargs["include"] = {"litellm_budget_table": True} - + response = await prisma_client.db.litellm_organizationtable.find_unique( **query_kwargs ) @@ -1844,7 +1845,9 @@ async def get_org_object( # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, - value=response.model_dump() if hasattr(response, "model_dump") else response, + value=response.model_dump() + if hasattr(response, "model_dump") + else response, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -2218,10 +2221,15 @@ async def _virtual_key_max_budget_alert_check( and valid_token.spend is not None and valid_token.spend > 0 ): - alert_threshold = valid_token.max_budget * EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE - + alert_threshold = ( + valid_token.max_budget * EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE + ) + # Only alert if we've crossed the threshold but haven't exceeded max_budget yet - if valid_token.spend >= alert_threshold and valid_token.spend < valid_token.max_budget: + if ( + valid_token.spend >= alert_threshold + and valid_token.spend < valid_token.max_budget + ): verbose_proxy_logger.debug( "Reached Max Budget Alert Threshold for token %s, spend %s, max_budget %s, alert_threshold %s", valid_token.token, @@ -2274,7 +2282,7 @@ async def _check_team_member_budget( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) - + if ( team_membership is not None and team_membership.litellm_budget_table is not None @@ -2282,8 +2290,8 @@ async def _check_team_member_budget( ): team_member_budget = team_membership.litellm_budget_table.max_budget team_member_spend = team_membership.spend or 0.0 - - if team_member_spend > team_member_budget: + + if team_member_spend >= team_member_budget: raise litellm.BudgetExceededError( current_cost=team_member_spend, max_budget=team_member_budget, @@ -2343,11 +2351,11 @@ async def _organization_max_budget_check( ): """ Check if the organization is over its max budget. - + This function checks the organization budget using: 1. First, tries to use valid_token.org_id (if key has organization_id set) 2. Falls back to team_object.organization_id (if key doesn't have org_id but team does) - + This ensures organization budget checks work even when keys don't have organization_id set directly, as long as their team belongs to an organization. @@ -2364,7 +2372,7 @@ async def _organization_max_budget_check( org_id = valid_token.org_id elif team_object is not None and team_object.organization_id is not None: org_id = team_object.organization_id - + # If no organization_id found, skip the check if org_id is None: return @@ -2655,4 +2663,4 @@ def _can_object_call_vector_stores( code=status.HTTP_401_UNAUTHORIZED, ) - return True \ No newline at end of file + return True From 37c014c80551825179d55ce1fe1d90602efd0fc7 Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Sat, 17 Jan 2026 00:17:20 +0530 Subject: [PATCH 3/9] ci(github): add automated duplicate issue checker and template safeguards (#19218) --- .github/ISSUE_TEMPLATE/bug_report.yml | 8 ++++++ .github/ISSUE_TEMPLATE/feature_request.yml | 8 ++++++ .github/workflows/check_duplicate_issues.yml | 29 ++++++++++++++++++++ 3 files changed, 45 insertions(+) create mode 100644 .github/workflows/check_duplicate_issues.yml diff --git a/.github/ISSUE_TEMPLATE/bug_report.yml b/.github/ISSUE_TEMPLATE/bug_report.yml index e0c1051dd29..bbe4b76775d 100644 --- a/.github/ISSUE_TEMPLATE/bug_report.yml +++ b/.github/ISSUE_TEMPLATE/bug_report.yml @@ -9,6 +9,14 @@ body: Thanks for taking the time to fill out this bug report! **💡 Tip:** See our [Troubleshooting Guide](https://docs.litellm.ai/docs/troubleshoot) for what information to include. + - type: checkboxes + id: duplicate-check + attributes: + label: Check for existing issues + description: Please search to see if an issue already exists for the bug you encountered. + options: + - label: I have searched the existing issues and checked that my issue is not a duplicate. + required: true - type: textarea id: what-happened attributes: diff --git a/.github/ISSUE_TEMPLATE/feature_request.yml b/.github/ISSUE_TEMPLATE/feature_request.yml index e575db7302a..4cc42901897 100644 --- a/.github/ISSUE_TEMPLATE/feature_request.yml +++ b/.github/ISSUE_TEMPLATE/feature_request.yml @@ -7,6 +7,14 @@ body: attributes: value: | Thanks for making LiteLLM better! + - type: checkboxes + id: duplicate-check + attributes: + label: Check for existing issues + description: Please search to see if an issue already exists for the feature you are requesting. + options: + - label: I have searched the existing issues and checked that my issue is not a duplicate. + required: true - type: textarea id: the-feature attributes: diff --git a/.github/workflows/check_duplicate_issues.yml b/.github/workflows/check_duplicate_issues.yml new file mode 100644 index 00000000000..14d6964fcdb --- /dev/null +++ b/.github/workflows/check_duplicate_issues.yml @@ -0,0 +1,29 @@ +name: Check Duplicate Issues + +on: + issues: + types: [opened, edited] + +jobs: + check-duplicate: + runs-on: ubuntu-latest + permissions: + issues: write + contents: read + steps: + - name: Check for potential duplicates + uses: wow-actions/potential-duplicates@v1 + with: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + label: potential-duplicate + threshold: 0.6 + reaction: eyes + comment: | + **⚠️ Potential duplicate detected** + + This issue appears similar to existing issue(s): + {{#issues}} + - [#{{number}}]({{html_url}}) - {{title}} ({{accuracy}}% similar) + {{/issues}} + + Please review the linked issue(s) to see if they address your concern. If this is not a duplicate, please provide additional context to help us understand the difference. From 2c75194b491e66a7cc7f5693dbee65cea6b751da Mon Sep 17 00:00:00 2001 From: Anand Kamble Date: Fri, 16 Jan 2026 11:26:15 -0800 Subject: [PATCH 4/9] fix(vertex_ai): Vertex AI 400 Error: Model used by GenerateContent request (models/gemini-3-*) and CachedContent (models/gemini-3-*) has to be the same (#19193) * fix(vertex_ai): include model in context cache key generation * test(vertex_ai): update context caching tests to verify model in cache key --- .../context_caching/vertex_ai_context_caching.py | 4 ++-- .../context_caching/test_vertex_ai_context_caching.py | 8 ++++---- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index cff1bebceb9..289963e917a 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -304,7 +304,7 @@ class ContextCachingEndpoints(VertexBase): ## CHECK IF CACHED ALREADY generated_cache_key = local_cache_obj.get_cache_key( - messages=cached_messages, tools=tools + messages=cached_messages, tools=tools, model=model ) google_cache_name = self.check_cache( cache_key=generated_cache_key, @@ -433,7 +433,7 @@ class ContextCachingEndpoints(VertexBase): ## CHECK IF CACHED ALREADY generated_cache_key = local_cache_obj.get_cache_key( - messages=cached_messages, tools=tools + messages=cached_messages, tools=tools, model=model ) google_cache_name = await self.async_check_cache( cache_key=generated_cache_key, diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 88d1b59c5b5..e9d14d4e18f 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -187,9 +187,9 @@ class TestContextCachingEndpoints: assert returned_params == optional_params assert returned_cache == "existing_cache_name" - # Verify cache key was generated with tools + # Verify cache key was generated with tools and model mock_cache_obj.get_cache_key.assert_called_once_with( - messages=cached_messages, tools=self.sample_tools + messages=cached_messages, tools=self.sample_tools, model="gemini-1.5-pro" ) @pytest.mark.parametrize( @@ -460,9 +460,9 @@ class TestContextCachingEndpoints: assert returned_params == optional_params assert returned_cache == "existing_cache_name" - # Verify cache key was generated with tools + # Verify cache key was generated with tools and model mock_cache_obj.get_cache_key.assert_called_once_with( - messages=cached_messages, tools=self.sample_tools + messages=cached_messages, tools=self.sample_tools, model="gemini-1.5-pro" ) @pytest.mark.asyncio From 17f8916ce3d54b0bb44c1581ed72afc0f9b6f5c5 Mon Sep 17 00:00:00 2001 From: Graham Neubig Date: Fri, 16 Jan 2026 14:29:38 -0500 Subject: [PATCH 5/9] fix(logging): Include langfuse logger in JSON logging when langfuse callback is used (#19162) When JSON_LOGS is enabled and langfuse is configured as a success/failure callback, the langfuse logger now receives the JSON formatter. This ensures langfuse SDK log messages (like 'Item exceeds size limit' warnings) are output as JSON with proper level information, instead of plain text that log aggregators may incorrectly classify as errors. Fixes issue where langfuse warnings appeared as errors in Datadog due to missing log level in unformatted output. Co-authored-by: openhands --- litellm/_logging.py | 22 +++++++++++++++++++++- 1 file changed, 21 insertions(+), 1 deletion(-) diff --git a/litellm/_logging.py b/litellm/_logging.py index 73902d2fc5a..b3156b15ba7 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -133,6 +133,26 @@ ALL_LOGGERS = [ ] +def _get_loggers_to_initialize(): + """ + Get all loggers that should be initialized with the JSON handler. + + Includes third-party integration loggers (like langfuse) if they are + configured as callbacks. + """ + import litellm + + loggers = list(ALL_LOGGERS) + + # Add langfuse logger if langfuse is being used as a callback + langfuse_callbacks = {"langfuse", "langfuse_otel"} + all_callbacks = set(litellm.success_callback + litellm.failure_callback) + if langfuse_callbacks & all_callbacks: + loggers.append(logging.getLogger("langfuse")) + + return loggers + + def _initialize_loggers_with_handler(handler: logging.Handler): """ Initialize all loggers with a handler @@ -140,7 +160,7 @@ def _initialize_loggers_with_handler(handler: logging.Handler): - Adds a handler to each logger - Prevents bubbling to parent/root (critical to prevent duplicate JSON logs) """ - for lg in ALL_LOGGERS: + for lg in _get_loggers_to_initialize(): lg.handlers.clear() # remove any existing handlers lg.addHandler(handler) # add JSON formatter handler lg.propagate = False # prevent bubbling to parent/root From 237ba2203ec619721c323c5b6ab471444fb4f78b Mon Sep 17 00:00:00 2001 From: YutaSaito <36355491+uc4w6c@users.noreply.github.com> Date: Sat, 17 Jan 2026 05:57:07 +0900 Subject: [PATCH 6/9] Revert "[Fix] /user/new Privilege Escalation" --- .../internal_user_endpoints.py | 7 -- .../test_internal_user_endpoints.py | 83 ------------------- 2 files changed, 90 deletions(-) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 89ecc31d83b..1850ffa2560 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -412,13 +412,6 @@ async def new_user( status_code=403, detail="License is over limit. Please contact support@berri.ai to upgrade your license.", ) - - # Only proxy admins can create administrative users - if data.user_role in [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY] and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException( - status_code=403, - detail=f"Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). Attempted to create user with role: {data.user_role}. Your role: {user_api_key_dict.user_role}" - ) data_json = data.json() # type: ignore data_json = _update_internal_new_user_params(data_json, data) 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 397a6af556f..33f2a75fac6 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 @@ -12,7 +12,6 @@ sys.path.insert( from litellm.proxy._types import ( LiteLLM_UserTableFiltered, - LitellmUserRoles, NewUserRequest, ProxyException, UpdateUserRequest, @@ -307,88 +306,6 @@ async def test_new_user_license_over_limit(mocker): mock_license_check.is_over_limit.assert_called_once_with(total_users=1000) -@pytest.mark.asyncio -async def test_new_user_non_admin_cannot_create_admin(mocker): - """ - Test that non-admin users cannot create administrative users (PROXY_ADMIN or PROXY_ADMIN_VIEW_ONLY). - This prevents privilege escalation vulnerabilities. - """ - from litellm.proxy.management_endpoints.internal_user_endpoints import new_user - - # Mock the prisma client - mock_prisma_client = mocker.MagicMock() - - # Setup the mock count response (under license limit) - async def mock_count(*args, **kwargs): - return 5 # Low user count, under limit - - mock_prisma_client.db.litellm_usertable.count = mock_count - - # Mock duplicate checks to pass - async def mock_check_duplicate_user_email(*args, **kwargs): - return None # No duplicate found - - async def mock_check_duplicate_user_id(*args, **kwargs): - return None # No duplicate found - - mocker.patch( - "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email", - mock_check_duplicate_user_email, - ) - mocker.patch( - "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id", - mock_check_duplicate_user_id, - ) - - # Mock the license check to return False (under limit) - mock_license_check = mocker.MagicMock() - mock_license_check.is_over_limit.return_value = False - - # Patch the imports in the endpoint - mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check) - - # Test Case 1: INTERNAL_USER trying to create PROXY_ADMIN - user_request = NewUserRequest( - user_email="admin@example.com", user_role=LitellmUserRoles.PROXY_ADMIN - ) - - # Mock user_api_key_dict with non-admin role - mock_user_api_key_dict = UserAPIKeyAuth( - user_id="test_internal_user", user_role=LitellmUserRoles.INTERNAL_USER - ) - - # Call new_user function and expect ProxyException - with pytest.raises(ProxyException) as exc_info: - await new_user(data=user_request, user_api_key_dict=mock_user_api_key_dict) - - # Verify the exception details - assert exc_info.value.code == 403 or exc_info.value.code == "403" - assert "Only proxy admins can create administrative users" in str(exc_info.value.message) - assert "proxy_admin" in str(exc_info.value.message) - assert "proxy_admin_viewer" in str(exc_info.value.message) - assert str(LitellmUserRoles.PROXY_ADMIN) in str(exc_info.value.message) - assert str(LitellmUserRoles.INTERNAL_USER) in str(exc_info.value.message) - - # Test Case 2: INTERNAL_USER trying to create PROXY_ADMIN_VIEW_ONLY - user_request_viewer = NewUserRequest( - user_email="admin_viewer@example.com", - user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, - ) - - with pytest.raises(ProxyException) as exc_info2: - await new_user( - data=user_request_viewer, user_api_key_dict=mock_user_api_key_dict - ) - - # Verify the exception details - assert exc_info2.value.code == 403 or exc_info2.value.code == "403" - assert "Only proxy admins can create administrative users" in str( - exc_info2.value.message - ) - assert str(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) in str(exc_info2.value.message) - - @pytest.mark.asyncio async def test_user_info_url_encoding_plus_character(mocker): """ From 66d67ae3563dbe31335aa06001478743afd80789 Mon Sep 17 00:00:00 2001 From: YutaSaito <36355491+uc4w6c@users.noreply.github.com> Date: Sat, 17 Jan 2026 06:01:12 +0900 Subject: [PATCH 7/9] Revert "Add sanititzation for anthropic messages" --- .../docs/completion/message_sanitization.md | 468 ------------------ docs/my-website/sidebars.js | 1 - .../prompt_templates/factory.py | 220 -------- .../anthropic/test_message_sanitization.py | 380 -------------- 4 files changed, 1069 deletions(-) delete mode 100644 docs/my-website/docs/completion/message_sanitization.md delete mode 100644 tests/test_litellm/llms/anthropic/test_message_sanitization.py diff --git a/docs/my-website/docs/completion/message_sanitization.md b/docs/my-website/docs/completion/message_sanitization.md deleted file mode 100644 index 0a1f766e2fd..00000000000 --- a/docs/my-website/docs/completion/message_sanitization.md +++ /dev/null @@ -1,468 +0,0 @@ -import Tabs from '@theme/Tabs'; -import TabItem from '@theme/TabItem'; - -# Message Sanitization for Tool Calling for anthropic models - -**Automatically fix common message formatting issues when using tool calling with `modify_params=True`** - -LiteLLM can automatically sanitize messages to handle common issues that occur during tool calling workflows, especially when using OpenAI-compatible clients with providers that have strict message format requirements (like Anthropic Claude). - -## Overview - -When `litellm.modify_params = True` is enabled, LiteLLM automatically sanitizes messages to fix three common issues: - -1. **Orphaned Tool Calls** - Assistant messages with tool_calls but missing tool results -2. **Orphaned Tool Results** - Tool messages that reference non-existent tool_call_ids -3. **Empty Message Content** - Messages with empty or whitespace-only text content - -This ensures your tool calling workflows work seamlessly across different LLM providers without manual message validation. - -## Why Message Sanitization? - -Different LLM providers have varying requirements for message formats, especially during tool calling: - -- **Anthropic Claude** requires every tool_call to have a corresponding tool result -- Some providers reject messages with empty content -- OpenAI-compatible clients may not always maintain perfect message consistency - -Without sanitization, these issues cause API errors that interrupt your workflows. With `modify_params=True`, LiteLLM handles these edge cases automatically. - -## Quick Start - - - - -```python -import litellm - -# Enable automatic message sanitization -litellm.modify_params = True - -# This will work even if messages have formatting issues -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=[ - {"role": "user", "content": "What's the weather in Boston?"}, - { - "role": "assistant", - "tool_calls": [ - { - "id": "call_123", - "type": "function", - "function": {"name": "get_weather", "arguments": '{"city": "Boston"}'} - } - ] - # Missing tool result - LiteLLM will add a dummy result automatically - }, - {"role": "user", "content": "Thanks!"} - ], - tools=[{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a city", - "parameters": { - "type": "object", - "properties": {"city": {"type": "string"}}, - "required": ["city"] - } - } - }] -) -``` - - - - -```yaml -litellm_settings: - modify_params: true # Enable automatic message sanitization - -model_list: - - model_name: claude-3-5-sonnet - litellm_params: - model: anthropic/claude-3-5-sonnet-20241022 -``` - - - - -## Sanitization Cases - -### Case A: Orphaned Tool Calls (Missing Tool Results) - -**Problem:** An assistant message contains `tool_calls`, but no corresponding tool result messages follow. - -**Solution:** LiteLLM automatically adds dummy tool result messages for any missing tool results. - -**Example:** - -```python -import litellm -litellm.modify_params = True - -# Messages with orphaned tool calls -messages = [ - {"role": "user", "content": "Search for Python tutorials"}, - { - "role": "assistant", - "tool_calls": [ - { - "id": "call_abc123", - "type": "function", - "function": {"name": "web_search", "arguments": '{"query": "Python tutorials"}'} - } - ] - }, - # Missing tool result here! - {"role": "user", "content": "What about JavaScript?"} -] - -# LiteLLM automatically adds: -# { -# "role": "tool", -# "tool_call_id": "call_abc123", -# "content": "[System: Tool execution skipped/interrupted by user. No result provided for tool 'web_search'.]" -# } - -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=messages, - tools=[...] -) -``` - -**When this happens:** -- User interrupts tool execution -- Client loses tool results due to network issues -- Conversation flow changes before tool completes -- Multi-turn conversations where tools are optional - -### Case B: Orphaned Tool Results (Invalid tool_call_id) - -**Problem:** A tool message references a `tool_call_id` that doesn't exist in any previous assistant message. - -**Solution:** LiteLLM automatically removes these orphaned tool result messages. - -**Example:** - -```python -import litellm -litellm.modify_params = True - -# Messages with orphaned tool result -messages = [ - {"role": "user", "content": "Hello"}, - {"role": "assistant", "content": "Hi! How can I help?"}, - { - "role": "tool", - "tool_call_id": "call_nonexistent", # This tool_call_id doesn't exist! - "content": "Some result" - } -] - -# LiteLLM automatically removes the orphaned tool message - -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=messages -) -``` - -**When this happens:** -- Message history is manually edited -- Tool results are duplicated or mismatched -- Conversation state is restored incorrectly -- Messages are merged from different conversations - -### Case C: Empty Message Content - -**Problem:** User or assistant messages have empty or whitespace-only content. - -**Solution:** LiteLLM replaces empty content with a system placeholder message. - -**Example:** - -```python -import litellm -litellm.modify_params = True - -# Messages with empty content -messages = [ - {"role": "user", "content": ""}, # Empty content - {"role": "assistant", "content": " "}, # Whitespace only -] - -# LiteLLM automatically replaces with: -# {"role": "user", "content": "[System: Empty message content sanitised to satisfy protocol]"} -# {"role": "assistant", "content": "[System: Empty message content sanitised to satisfy protocol]"} - -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=messages -) -``` - -**When this happens:** -- UI sends empty messages -- Content is stripped during preprocessing -- Placeholder messages in conversation history -- Edge cases in message construction - -## Configuration - -### Enable Globally - - - - -```python -import litellm - -# Enable for all completion calls -litellm.modify_params = True -``` - - - - -```yaml -litellm_settings: - modify_params: true -``` - - - - -```bash -export LITELLM_MODIFY_PARAMS=True -``` - - - - -### Enable Per-Request - -```python -import litellm - -# Enable only for specific requests -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=messages, - modify_params=True # Override global setting -) -``` - -## Supported Providers - -Message sanitization works with all LLM providers that support tool calling: - -- ✅ Anthropic (Claude) -- ✅ OpenAI (GPT-4, GPT-3.5) -- ✅ AWS Bedrock (Claude, Titan) -- ✅ Google Vertex AI (Claude, Gemini) -- ✅ Azure OpenAI -- ✅ And all other providers with tool calling support - -## Implementation Details - -### How It Works - -The message sanitization process runs **before** messages are converted to provider-specific formats: - -1. **Input:** OpenAI-format messages with potential issues -2. **Sanitization:** Three helper functions process the messages: - - `_sanitize_empty_text_content()` - Fixes empty content - - `_add_missing_tool_results()` - Adds dummy tool results - - `_is_orphaned_tool_result()` - Identifies orphaned results -3. **Output:** Clean, provider-compatible messages - -### Code Reference - -The sanitization logic is implemented in: -- `litellm/litellm_core_utils/prompt_templates/factory.py` -- Function: `sanitize_messages_for_tool_calling()` - -### Logging - -When sanitization occurs, LiteLLM logs debug messages: - -```python -import litellm -litellm.set_verbose = True # Enable debug logging - -# You'll see logs like: -# "_add_missing_tool_results: Found 1 orphaned tool calls. Adding dummy tool results." -# "_is_orphaned_tool_result: Found orphaned tool result with tool_call_id=call_123" -# "_sanitize_empty_text_content: Replaced empty text content in user message" -``` - -## Best Practices - -### 1. Enable for Production Workflows - -```python -# Recommended for production -litellm.modify_params = True - -# Ensures robust handling of edge cases -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=messages, - tools=tools -) -``` - -### 2. Preserve Tool Results When Possible - -While sanitization handles missing tool results, it's better to provide actual results: - -```python -# Good: Provide actual tool results -messages = [ - {"role": "user", "content": "Search for Python"}, - {"role": "assistant", "tool_calls": [...]}, - {"role": "tool", "tool_call_id": "call_123", "content": "Actual search results"} -] - -# Fallback: Sanitization adds dummy result if missing -messages = [ - {"role": "user", "content": "Search for Python"}, - {"role": "assistant", "tool_calls": [...]}, - # Missing tool result - sanitization adds dummy -] -``` - -### 3. Monitor Sanitization Events - -Use logging to track when sanitization occurs: - -```python -import litellm -import logging - -# Enable debug logging -litellm.set_verbose = True -logging.basicConfig(level=logging.DEBUG) - -# Track sanitization events in your application -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=messages -) -``` - -### 4. Test Edge Cases - -Ensure your application handles sanitized messages correctly: - -```python -import litellm -litellm.modify_params = True - -# Test orphaned tool calls -test_messages = [ - {"role": "user", "content": "Test"}, - {"role": "assistant", "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "test", "arguments": "{}"}}]}, - {"role": "user", "content": "Continue"} # No tool result -] - -response = litellm.completion( - model="anthropic/claude-3-5-sonnet-20241022", - messages=test_messages, - tools=[...] -) - -# Verify the response handles the dummy tool result appropriately -``` - -## Related Features - -- **[Drop Params](./drop_params.md)** - Drop unsupported parameters for specific providers -- **[Message Trimming](./message_trimming.md)** - Trim messages to fit token limits -- **[Function Calling](./function_call.md)** - Complete guide to tool/function calling -- **[Reasoning Content](../reasoning_content.md)** - Extended thinking with tool calling - -## Troubleshooting - -### Sanitization Not Working - -**Issue:** Messages still cause errors despite `modify_params=True` - -**Solution:** -1. Verify `modify_params` is enabled: - ```python - import litellm - print(litellm.modify_params) # Should be True - ``` - -2. Check if the issue is provider-specific: - ```python - litellm.set_verbose = True # Enable debug logging - ``` - -3. Ensure you're using a recent version of LiteLLM: - ```bash - pip install --upgrade litellm - ``` - -### Unexpected Dummy Tool Results - -**Issue:** Dummy tool results appear when you expect actual results - -**Cause:** Tool result messages are missing or have incorrect `tool_call_id` - -**Solution:** -1. Verify tool result messages have correct `tool_call_id`: - ```python - # Correct - {"role": "tool", "tool_call_id": "call_123", "content": "result"} - - # Incorrect - will be treated as orphaned - {"role": "tool", "tool_call_id": "wrong_id", "content": "result"} - ``` - -2. Ensure tool results immediately follow assistant messages with tool_calls - -### Performance Impact - -**Issue:** Concerned about performance overhead - -**Details:** Message sanitization has minimal performance impact: -- Runs in O(n) time where n = number of messages -- Only processes messages when `modify_params=True` -- Typically adds < 1ms to request processing time - -## FAQ - -**Q: Does sanitization modify my original messages?** - -A: No, sanitization creates a new list of messages. Your original messages remain unchanged. - -**Q: Can I disable specific sanitization cases?** - -A: Currently, all three cases are handled together when `modify_params=True`. To disable sanitization entirely, set `modify_params=False`. - -**Q: What happens to the dummy tool results?** - -A: Dummy tool results are sent to the LLM provider along with other messages. The model sees them as regular tool results with informative error messages. - -**Q: Does this work with streaming?** - -A: Yes, message sanitization works with both streaming and non-streaming requests. - -**Q: Is this related to `drop_params`?** - -A: No, they're separate features: -- `modify_params` - Modifies/fixes message content and structure -- `drop_params` - Removes unsupported API parameters - -Both can be enabled simultaneously. - -## See Also - -- [Reasoning Content with Tool Calling](../reasoning_content.md) -- [Function Calling Guide](./function_call.md) -- [Bedrock Provider Documentation](../providers/bedrock.md) -- [Anthropic Provider Documentation](../providers/anthropic.md) diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index acc5d538550..38a26f6b183 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -822,7 +822,6 @@ const sidebars = { "completion/knowledgebase", "guides/code_interpreter", "completion/message_trimming", - "completion/message_sanitization", "completion/model_alias", "completion/mock_requests", "completion/predict_outputs", diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 2311b34a2cc..01bf18d79b2 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1989,223 +1989,6 @@ def anthropic_process_openai_file_message( ) -def _sanitize_empty_text_content( - message: AllMessageValues, -) -> AllMessageValues: - """ - Case C: Sanitize empty text content - - Replace empty or whitespace-only text content with a placeholder message. - - Returns: - The message with sanitized content if needed, otherwise the original message - """ - if message.get("role") in ["user", "assistant"]: - content = message.get("content") - if isinstance(content, str): - if not content or not content.strip(): - message = dict(message) # Make a copy - message["content"] = "[System: Empty message content sanitised to satisfy protocol]" - verbose_logger.debug( - f"_sanitize_empty_text_content: Replaced empty text content in {message.get('role')} message" - ) - return message - - -def _add_missing_tool_results( - current_message: AllMessageValues, - messages: List[AllMessageValues], - current_index: int, -) -> List[AllMessageValues]: - """ - Case A: Missing tool_result for tool_use (orphaned tool calls) - - If an assistant message has tool_calls but no corresponding tool result follows, - add a dummy tool result message indicating the user did not provide the result. - - Returns: - A list containing the assistant message followed by any dummy tool results needed - """ - result_messages: List[AllMessageValues] = [] - tool_calls = current_message.get("tool_calls") - - if not tool_calls or len(tool_calls) == 0: - return [current_message] - - # Collect all tool_call_ids from this assistant message - expected_tool_call_ids = set() - for tool_call in tool_calls: - tool_call_id = None - if isinstance(tool_call, dict): - tool_call_id = tool_call.get("id") - else: - tool_call_id = getattr(tool_call, "id", None) - if tool_call_id: - expected_tool_call_ids.add(tool_call_id) - - found_tool_call_ids = set() - j = current_index + 1 - - while j < len(messages): - next_msg = messages[j] - next_role = next_msg.get("role") - - if next_role == "assistant": - break - - if next_role in ["tool", "function"]: - tool_call_id = next_msg.get("tool_call_id") - if tool_call_id: - found_tool_call_ids.add(tool_call_id) - - j += 1 - - # Find missing tool results - missing_tool_call_ids = expected_tool_call_ids - found_tool_call_ids - - if missing_tool_call_ids: - verbose_logger.debug( - f"_add_missing_tool_results: Found {len(missing_tool_call_ids)} orphaned tool calls. Adding dummy tool results." - ) - - result_messages.append(current_message) - - for tool_call_id in missing_tool_call_ids: - tool_name = "unknown_tool" - for tool_call in tool_calls: - tc_id = None - if isinstance(tool_call, dict): - tc_id = tool_call.get("id") - else: - tc_id = getattr(tool_call, "id", None) - - if tc_id == tool_call_id: - if isinstance(tool_call, dict): - function = tool_call.get("function", {}) - if isinstance(function, dict): - tool_name = function.get("name", "unknown_tool") - else: - tool_name = getattr(function, "name", "unknown_tool") - else: - function = getattr(tool_call, "function", None) - if function: - tool_name = getattr(function, "name", "unknown_tool") - break - - dummy_tool_result: ChatCompletionToolMessage = { - "role": "tool", - "tool_call_id": tool_call_id, - "content": f"[System: Tool execution skipped/interrupted by user. No result provided for tool '{tool_name}'.]", - } - result_messages.append(dummy_tool_result) - - return result_messages - - return [current_message] - - -def _is_orphaned_tool_result( - current_message: AllMessageValues, - sanitized_messages: List[AllMessageValues], -) -> bool: - """ - Case B: Orphaned tool_result (unexpected result) - - Check if a tool message references a tool_call_id that doesn't exist in the previous - assistant message. - - Returns: - True if this is an orphaned tool result that should be removed, False otherwise - """ - if current_message.get("role") not in ["tool", "function"]: - return False - - tool_call_id = current_message.get("tool_call_id") - - if not tool_call_id: - return False - - # Look back to find the most recent assistant message with tool_calls - found_matching_tool_call = False - - for j in range(len(sanitized_messages) - 1, -1, -1): - prev_msg = sanitized_messages[j] - if prev_msg.get("role") == "assistant": - tool_calls = prev_msg.get("tool_calls") - if tool_calls: - for tool_call in tool_calls: - tc_id = None - if isinstance(tool_call, dict): - tc_id = tool_call.get("id") - else: - tc_id = getattr(tool_call, "id", None) - - if tc_id == tool_call_id: - found_matching_tool_call = True - break - - break - - if not found_matching_tool_call: - verbose_logger.debug( - "_is_orphaned_tool_result: Found orphaned tool result with redacted tool_call_id" - ) - return True - - return False - - -def sanitize_messages_for_tool_calling( - messages: List[AllMessageValues], -) -> List[AllMessageValues]: - """ - Sanitize messages for tool calling to handle common issues when modify_params=True: - - Case A: Missing tool_result for tool_use (orphaned tool calls) - - If an assistant message has tool_calls but no corresponding tool result follows, - add a dummy tool result message indicating the user did not provide the result. - - Case B: Orphaned tool_result (unexpected result) - - If a tool message references a tool_call_id that doesn't exist in the previous - assistant message, remove that tool message. - - Case C: Empty text content - - Replace empty or whitespace-only text content with a placeholder message. - - This function operates on OpenAI format messages before they are converted to - provider-specific formats. - """ - if not litellm.modify_params: - return messages - - sanitized_messages: List[AllMessageValues] = [] - i = 0 - - while i < len(messages): - current_message = messages[i] - - # Case C: Sanitize empty text content - current_message = _sanitize_empty_text_content(current_message) - - # Case A: Check if assistant message has tool_calls without following tool results - if current_message.get("role") == "assistant": - result_messages = _add_missing_tool_results(current_message, messages, i) - - # If dummy tool results were added, extend sanitized_messages and continue - if len(result_messages) > 1: - sanitized_messages.extend(result_messages) - i += 1 - continue - - # Case B: Check for orphaned tool results - if _is_orphaned_tool_result(current_message, sanitized_messages): - i += 1 - continue # Skip this orphaned tool result - - # Add the message to sanitized list - sanitized_messages.append(current_message) - i += 1 - - return sanitized_messages - - def anthropic_messages_pt( # noqa: PLR0915 messages: List[AllMessageValues], model: str, @@ -2225,9 +2008,6 @@ def anthropic_messages_pt( # noqa: PLR0915 5. System messages are a separate param to the Messages API 6. Ensure we only accept role, content. (message.name is not supported) """ - # Sanitize messages for tool calling issues when modify_params=True - messages = sanitize_messages_for_tool_calling(messages) - # add role=tool support to allow function call result/error submission user_message_types = {"user", "tool", "function"} # reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, merge them. diff --git a/tests/test_litellm/llms/anthropic/test_message_sanitization.py b/tests/test_litellm/llms/anthropic/test_message_sanitization.py deleted file mode 100644 index 489ef527b48..00000000000 --- a/tests/test_litellm/llms/anthropic/test_message_sanitization.py +++ /dev/null @@ -1,380 +0,0 @@ -""" -Test message sanitization for Anthropic API when modify_params=True - -Tests three cases: -A. Missing tool_result for tool_use (orphaned tool calls) -B. Orphaned tool_result without matching tool_use -C. Empty text content -""" - -import pytest -import sys -import os - -# Add the parent directory to the path so we can import litellm -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))) - -import litellm -from litellm.litellm_core_utils.prompt_templates.factory import ( - sanitize_messages_for_tool_calling, - anthropic_messages_pt, -) - - -class TestMessageSanitization: - """Test message sanitization for tool calling scenarios""" - - def setup_method(self): - """Setup for each test""" - # Save original modify_params value - self.original_modify_params = litellm.modify_params - litellm.modify_params = True - - def teardown_method(self): - """Cleanup after each test""" - # Restore original modify_params value - litellm.modify_params = self.original_modify_params - - def test_case_a_orphaned_tool_call_single(self): - """ - Test Case A: Assistant message with tool_calls but no tool result - Should add a dummy tool result message - """ - messages = [ - { - "role": "user", - "content": "What is the weather in Nashik?" - }, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "toolu_01Kus2cC3ydjBW7UK4GJqBP4", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "Nashik, India"}' - } - } - ] - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - # Should have 3 messages: user, assistant, and dummy tool result - assert len(sanitized) == 3 - assert sanitized[0]["role"] == "user" - assert sanitized[1]["role"] == "assistant" - assert sanitized[2]["role"] == "tool" - assert sanitized[2]["tool_call_id"] == "toolu_01Kus2cC3ydjBW7UK4GJqBP4" - assert "skipped" in sanitized[2]["content"].lower() or "interrupted" in sanitized[2]["content"].lower() - assert "get_weather" in sanitized[2]["content"] - - def test_case_a_orphaned_tool_call_multiple(self): - """ - Test Case A: Assistant message with multiple tool_calls, some missing results - """ - messages = [ - { - "role": "user", - "content": "Get weather for Nashik and Mumbai" - }, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "Nashik"}' - } - }, - { - "id": "call_2", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "Mumbai"}' - } - } - ] - }, - { - "role": "tool", - "tool_call_id": "call_1", - "content": "Weather in Nashik: 25°C" - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - # Should have 4 messages: user, assistant, tool result for call_1, dummy for call_2 - assert len(sanitized) == 4 - assert sanitized[0]["role"] == "user" - assert sanitized[1]["role"] == "assistant" - assert sanitized[2]["tool_call_id"] == "call_2" # Dummy added first - assert sanitized[3]["tool_call_id"] == "call_1" # Original tool result - - def test_case_b_orphaned_tool_result(self): - """ - Test Case B: Tool result without matching tool_call in previous assistant message - Should remove the orphaned tool result - """ - messages = [ - { - "role": "user", - "content": "Hello" - }, - { - "role": "assistant", - "content": "Hi there!" - }, - { - "role": "tool", - "tool_call_id": "nonexistent_id", - "content": "Some result" - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - # Should have only 2 messages, orphaned tool result removed - assert len(sanitized) == 2 - assert sanitized[0]["role"] == "user" - assert sanitized[1]["role"] == "assistant" - - def test_case_b_valid_tool_result_preserved(self): - """ - Test Case B: Valid tool result with matching tool_call should be preserved - """ - messages = [ - { - "role": "user", - "content": "What's the weather?" - }, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "Boston"}' - } - } - ] - }, - { - "role": "tool", - "tool_call_id": "call_123", - "content": "Weather: 20°C" - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - # All messages should be preserved - assert len(sanitized) == 3 - assert sanitized[2]["role"] == "tool" - assert sanitized[2]["tool_call_id"] == "call_123" - - def test_case_c_empty_text_content_user(self): - """ - Test Case C: Empty text content in user message - Should replace with placeholder - """ - messages = [ - { - "role": "user", - "content": "" - }, - { - "role": "assistant", - "content": "Hello!" - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - assert len(sanitized) == 2 - assert sanitized[0]["role"] == "user" - assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]" - - def test_case_c_whitespace_only_content(self): - """ - Test Case C: Whitespace-only content - Should replace with placeholder - """ - messages = [ - { - "role": "user", - "content": " \n \t " - }, - { - "role": "assistant", - "content": " " - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - assert len(sanitized) == 2 - assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]" - assert sanitized[1]["content"] == "[System: Empty message content sanitised to satisfy protocol]" - - def test_case_c_valid_content_preserved(self): - """ - Test Case C: Valid non-empty content should be preserved - """ - messages = [ - { - "role": "user", - "content": "Hello" - }, - { - "role": "assistant", - "content": "Hi there!" - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - assert len(sanitized) == 2 - assert sanitized[0]["content"] == "Hello" - assert sanitized[1]["content"] == "Hi there!" - - def test_combined_cases(self): - """ - Test combination of multiple cases - """ - messages = [ - { - "role": "user", - "content": "Get weather" - }, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "NYC"}' - } - } - ] - }, - # Missing tool result for call_1 - { - "role": "user", - "content": "" # Empty content - }, - { - "role": "assistant", - "content": "Response" - }, - { - "role": "tool", - "tool_call_id": "orphaned_id", # Orphaned tool result - "content": "Some data" - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - # Should have: user, assistant, dummy tool result, user (sanitized), assistant - # Orphaned tool result should be removed - assert len(sanitized) == 5 - assert sanitized[0]["role"] == "user" - assert sanitized[1]["role"] == "assistant" - assert sanitized[2]["role"] == "tool" - assert sanitized[2]["tool_call_id"] == "call_1" # Dummy added - assert sanitized[3]["role"] == "user" - assert sanitized[3]["content"] == "[System: Empty message content sanitised to satisfy protocol]" - assert sanitized[4]["role"] == "assistant" - - def test_modify_params_false_no_sanitization(self): - """ - Test that sanitization is skipped when modify_params=False - """ - litellm.modify_params = False - - messages = [ - { - "role": "user", - "content": "" - }, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{}' - } - } - ] - } - ] - - sanitized = sanitize_messages_for_tool_calling(messages) - - # Messages should be unchanged - assert len(sanitized) == 2 - assert sanitized[0]["content"] == "" - assert len(sanitized[1].get("tool_calls", [])) == 1 - - def test_anthropic_messages_pt_integration(self): - """ - Test that sanitization is integrated into anthropic_messages_pt - """ - litellm.modify_params = True - - messages = [ - { - "role": "user", - "content": "What is the weather in Nashik?" - }, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "toolu_01Kus2cC3ydjBW7UK4GJqBP4", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "Nashik, India"}' - } - } - ] - } - ] - - # This should not raise an error and should add dummy tool result - result = anthropic_messages_pt( - messages=messages, - model="claude-sonnet-4-5", - llm_provider="anthropic" - ) - - # Should have at least 2 messages (user and assistant) - # The tool result will be merged into user content - assert len(result) >= 2 - assert result[0]["role"] == "user" - assert result[1]["role"] == "assistant" - - -if __name__ == "__main__": - pytest.main([__file__, "-v"]) From ca2019776e9ecc387325495a030a4b5fcf57ceaf Mon Sep 17 00:00:00 2001 From: YutaSaito <36355491+uc4w6c@users.noreply.github.com> Date: Sat, 17 Jan 2026 06:04:24 +0900 Subject: [PATCH 8/9] Revert "Fix: malformed tool call transformation in bedrock" --- .../prompt_templates/factory.py | 20 +-- .../bedrock/chat/converse_transformation.py | 9 +- litellm/types/llms/bedrock.py | 2 +- .../test_bedrock_completion.py | 154 ------------------ 4 files changed, 10 insertions(+), 175 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 01bf18d79b2..4320f756454 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -3233,21 +3233,17 @@ def _convert_to_bedrock_tool_call_invoke( id = tool["id"] name = tool["function"].get("name", "") arguments = tool["function"].get("arguments", "") + arguments_dict = json.loads(arguments) if arguments else {} + # Ensure arguments_dict is always a dict (Bedrock requires toolUse.input to be an object) + # When some providers return arguments: '""' (JSON-encoded empty string), json.loads returns "" + if not isinstance(arguments_dict, dict): + arguments_dict = {} if not arguments or not arguments.strip(): - arguments_input = {} + arguments_dict = {} else: - # Try to parse the arguments JSON - try: - arguments_input = json.loads(arguments) - except json.JSONDecodeError as e: - verbose_logger.warning( - f"Malformed JSON in tool call arguments for tool '{name}': {str(e)}. " - f"Storing as raw string to allow conversation to continue." - ) - arguments_input = arguments - + arguments_dict = json.loads(arguments) bedrock_tool = BedrockToolUseBlock( - input=arguments_input, name=name, toolUseId=id + input=arguments_dict, name=name, toolUseId=id ) bedrock_content_block = BedrockContentBlock(toolUse=bedrock_tool) _parts_list.append(bedrock_content_block) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 9bc1e8c85e2..59590e464fc 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1395,16 +1395,9 @@ class AmazonConverseConfig(BaseConfig): response_tool_name = get_bedrock_tool_name( response_tool_name=_response_tool_name ) - tool_input = content["toolUse"]["input"] - if isinstance(tool_input, str): - arguments_str = tool_input - else: - # Otherwise, serialize it to JSON - arguments_str = json.dumps(tool_input) - _function_chunk = ChatCompletionToolCallFunctionChunk( name=response_tool_name, - arguments=arguments_str, + arguments=json.dumps(content["toolUse"]["input"]), ) _tool_response_chunk = ChatCompletionToolCallChunk( diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index e0858898eae..ef2f1ba4d5e 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -62,7 +62,7 @@ class ToolResultBlock(TypedDict, total=False): class ToolUseBlock(TypedDict): - input: Any # Per boto3 spec: document type can be dict, list, int, float, str, bool, or None + input: dict name: str toolUseId: str diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index f08060214c5..7c0db41d13a 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -3954,157 +3954,3 @@ def test_bedrock_openai_error_handling(): assert exc_info.value.status_code == 422 print("✓ Error handling works correctly") - - -def test_bedrock_malformed_tool_json_handling(): - """ - Test that Bedrock handles malformed JSON in tool call arguments gracefully. - - This test covers the issue where: - 1. LLM generates malformed JSON in tool call arguments - 2. Subsequent requests with conversation history should not crash - 3. The toolUse.input field should handle any JSON value type per boto3 spec - - Related issue: https://github.com/BerriAI/litellm/issues/[issue_number] - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - _convert_to_bedrock_tool_call_invoke, - ) - from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig - from litellm.types.llms.bedrock import ContentBlock - - # Test 1: Malformed JSON in tool call arguments - malformed_tool_calls = [ - { - "id": "call_123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "Paris", "invalid_json', # Malformed JSON - }, - } - ] - - # Should not raise an exception, but store as raw string - result = _convert_to_bedrock_tool_call_invoke(malformed_tool_calls) - assert len(result) == 1 - assert result[0]["toolUse"]["name"] == "get_weather" - # The malformed JSON should be stored as a string - assert isinstance(result[0]["toolUse"]["input"], str) - assert result[0]["toolUse"]["input"] == '{"location": "Paris", "invalid_json' - print("✓ Malformed JSON stored as raw string") - - # Test 2: Valid JSON should still work normally - valid_tool_calls = [ - { - "id": "call_456", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "London"}', - }, - } - ] - - result = _convert_to_bedrock_tool_call_invoke(valid_tool_calls) - assert len(result) == 1 - assert result[0]["toolUse"]["name"] == "get_weather" - assert isinstance(result[0]["toolUse"]["input"], dict) - assert result[0]["toolUse"]["input"] == {"location": "London"} - print("✓ Valid JSON parsed correctly") - - # Test 3: Empty arguments should create empty dict - empty_tool_calls = [ - { - "id": "call_789", - "type": "function", - "function": { - "name": "no_args_function", - "arguments": "", - }, - } - ] - - result = _convert_to_bedrock_tool_call_invoke(empty_tool_calls) - assert len(result) == 1 - assert result[0]["toolUse"]["input"] == {} - print("✓ Empty arguments handled correctly") - - # Test 4: Bedrock to OpenAI conversion handles string input - converse_config = AmazonConverseConfig() - content_blocks = [ - ContentBlock( - toolUse={ - "name": "get_weather", - "toolUseId": "call_123", - "input": '{"location": "Paris", "invalid_json', # String input (malformed) - } - ) - ] - - content_str, tools, reasoning = converse_config._translate_message_content( - content_blocks - ) - assert len(tools) == 1 - assert tools[0]["function"]["name"] == "get_weather" - # Should return the string as-is - assert tools[0]["function"]["arguments"] == '{"location": "Paris", "invalid_json' - print("✓ Bedrock to OpenAI conversion handles string input") - - # Test 5: Bedrock to OpenAI conversion handles dict input - content_blocks_dict = [ - ContentBlock( - toolUse={ - "name": "get_weather", - "toolUseId": "call_456", - "input": {"location": "London"}, # Dict input (normal case) - } - ) - ] - - content_str, tools, reasoning = converse_config._translate_message_content( - content_blocks_dict - ) - assert len(tools) == 1 - assert tools[0]["function"]["name"] == "get_weather" - # Should serialize dict to JSON string - assert tools[0]["function"]["arguments"] == '{"location": "London"}' - print("✓ Bedrock to OpenAI conversion handles dict input") - - # Test 6: Round-trip conversion with malformed JSON - # Test that we can convert OpenAI -> Bedrock -> OpenAI with malformed JSON - malformed_tool_calls_roundtrip = [ - { - "id": "call_999", - "type": "function", - "function": { - "name": "test_function", - "arguments": '{"key": "value", "broken', # Malformed - }, - } - ] - - # Step 1: OpenAI to Bedrock (should store as string) - bedrock_blocks = _convert_to_bedrock_tool_call_invoke(malformed_tool_calls_roundtrip) - assert isinstance(bedrock_blocks[0]["toolUse"]["input"], str) - - # Step 2: Bedrock back to OpenAI (should preserve the string) - content_blocks_roundtrip = [ - ContentBlock( - toolUse={ - "name": bedrock_blocks[0]["toolUse"]["name"], - "toolUseId": bedrock_blocks[0]["toolUse"]["toolUseId"], - "input": bedrock_blocks[0]["toolUse"]["input"], - } - ) - ] - - content_str, tools_roundtrip, reasoning = converse_config._translate_message_content( - content_blocks_roundtrip - ) - - # Should preserve the malformed JSON string through the round trip - assert tools_roundtrip[0]["function"]["arguments"] == '{"key": "value", "broken' - print("✓ Round-trip conversion preserves malformed JSON") - - print("✓ All malformed JSON handling tests passed") From bec61c39ae241e8ff60b04cf99a29a3ab6df6ca8 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Sat, 17 Jan 2026 06:17:38 +0900 Subject: [PATCH 9/9] =?UTF-8?q?bump:=20version=200.4.21=20=E2=86=92=200.4.?= =?UTF-8?q?22?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- litellm-proxy-extras/pyproject.toml | 4 ++-- pyproject.toml | 2 +- requirements.txt | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index 2952aa6c979..4304aaf9e96 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm-proxy-extras" -version = "0.4.21" +version = "0.4.22" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." authors = ["BerriAI"] readme = "README.md" @@ -22,7 +22,7 @@ requires = ["poetry-core"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "0.4.21" +version = "0.4.22" version_files = [ "pyproject.toml:version", "../requirements.txt:litellm-proxy-extras==", diff --git a/pyproject.toml b/pyproject.toml index a5071353d6b..55d97f9a98f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -59,7 +59,7 @@ websockets = {version = "^15.0.1", optional = true} boto3 = {version = "1.40.61", optional = true} redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"} mcp = {version = "^1.21.2", optional = true, python = ">=3.10"} -litellm-proxy-extras = {version = "0.4.21", optional = true} +litellm-proxy-extras = {version = "0.4.22", optional = true} rich = {version = "13.7.1", optional = true} litellm-enterprise = {version = "0.1.27", optional = true} diskcache = {version = "^5.6.1", optional = true} diff --git a/requirements.txt b/requirements.txt index e98e295de30..0880e04fc5f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -48,7 +48,7 @@ sentry_sdk==2.21.0 # for sentry error handling detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests cryptography==44.0.1 tzdata==2025.1 # IANA time zone database -litellm-proxy-extras==0.4.21 # for proxy extras - e.g. prisma migrations +litellm-proxy-extras==0.4.22 # for proxy extras - e.g. prisma migrations llm-sandbox==0.3.31 # for skill execution in sandbox ### LITELLM PACKAGE DEPENDENCIES python-dotenv==1.0.1 # for env