diff --git a/litellm/integrations/rubrik.py b/litellm/integrations/rubrik.py index eb68fa2db7d..5a4e020e8b1 100644 --- a/litellm/integrations/rubrik.py +++ b/litellm/integrations/rubrik.py @@ -6,6 +6,7 @@ import random import time import urllib.parse import uuid +from collections import Counter from typing import TYPE_CHECKING, Any, Literal, Optional import httpx @@ -513,12 +514,19 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger): returned_tool_calls = message.get("tool_calls") or [] blocking_explanation = message.get("content", "") - allowed_ids = { - tc["id"] for tc in returned_tool_calls if tc.get("id") is not None - } - allowed_count = sum(1 for tc in all_tool_calls if tc.id in allowed_ids) + allowed_id_counts: Counter = Counter( + tc["id"] + for tc in returned_tool_calls + if isinstance(tc, dict) and tc.get("id") + ) + required_id_counts: Counter = Counter(tc.id for tc in all_tool_calls if tc.id) - if allowed_count == len(all_tool_calls): + all_allowed = len(returned_tool_calls) >= len(all_tool_calls) and all( + allowed_id_counts.get(tc_id, 0) >= count + for tc_id, count in required_id_counts.items() + ) + + if all_allowed: return None explanation = blocking_explanation or "Tool call blocked by policy." diff --git a/litellm/responses/sse_output_recovery.py b/litellm/responses/sse_output_recovery.py index 70f470660ca..fb7b8e5d850 100644 --- a/litellm/responses/sse_output_recovery.py +++ b/litellm/responses/sse_output_recovery.py @@ -9,6 +9,8 @@ caller automatically applies to all of them. from typing import Any, Dict +_MAX_CONTENT_INDEX = 1024 + def record_output_item_chunk( parsed_chunk: Dict[str, Any], @@ -77,6 +79,9 @@ def record_output_text_chunk( except (TypeError, ValueError): content_index = len(content) + if content_index < 0 or content_index > _MAX_CONTENT_INDEX: + return + while len(content) <= content_index: content.append( { diff --git a/tests/test_litellm/integrations/test_rubrik.py b/tests/test_litellm/integrations/test_rubrik.py index 596f939d0c2..06234735e6d 100644 --- a/tests/test_litellm/integrations/test_rubrik.py +++ b/tests/test_litellm/integrations/test_rubrik.py @@ -864,6 +864,33 @@ class TestExtractBlockedTools: assert result is not None assert "blocked everything" in result + def test_duplicate_ids_block_when_only_one_returned(self): + from litellm.types.utils import ChatCompletionMessageToolCall, Function + + tc1 = ChatCompletionMessageToolCall( + id="call_dup", + type="function", + function=Function(name="fn", arguments="{}"), + ) + tc2 = ChatCompletionMessageToolCall( + id="call_dup", + type="function", + function=Function(name="fn", arguments="{}"), + ) + service_resp = { + "choices": [ + { + "message": { + "tool_calls": [{"id": "call_dup"}], + "content": "blocked duplicate", + } + } + ] + } + result = RubrikLogger._extract_blocked_tools(service_resp, [tc1, tc2]) + assert result is not None + assert "blocked duplicate" in result + # -- Sanitize proxy server request ------------------------------------------- diff --git a/tests/test_litellm/responses/test_sse_output_recovery.py b/tests/test_litellm/responses/test_sse_output_recovery.py new file mode 100644 index 00000000000..c8f3325a624 --- /dev/null +++ b/tests/test_litellm/responses/test_sse_output_recovery.py @@ -0,0 +1,57 @@ +"""Tests for litellm.responses.sse_output_recovery helpers.""" + +from litellm.responses.sse_output_recovery import ( + _MAX_CONTENT_INDEX, + record_output_text_chunk, +) + + +def test_text_chunk_with_oversized_content_index_is_dropped(): + output_items: dict = {} + text_only_items: dict = {} + record_output_text_chunk( + parsed_chunk={ + "type": "response.output_text.done", + "output_index": 0, + "content_index": _MAX_CONTENT_INDEX + 1, + "text": "ignored", + }, + output_items=output_items, + text_only_items=text_only_items, + ) + item = text_only_items[0] + assert item["content"] == [] + + +def test_text_chunk_with_negative_content_index_is_dropped(): + output_items: dict = {} + text_only_items: dict = {} + record_output_text_chunk( + parsed_chunk={ + "type": "response.output_text.done", + "output_index": 0, + "content_index": -1, + "text": "ignored", + }, + output_items=output_items, + text_only_items=text_only_items, + ) + assert text_only_items[0]["content"] == [] + + +def test_text_chunk_at_max_content_index_is_recorded(): + output_items: dict = {} + text_only_items: dict = {} + record_output_text_chunk( + parsed_chunk={ + "type": "response.output_text.done", + "output_index": 0, + "content_index": _MAX_CONTENT_INDEX, + "text": "kept", + }, + output_items=output_items, + text_only_items=text_only_items, + ) + content = text_only_items[0]["content"] + assert len(content) == _MAX_CONTENT_INDEX + 1 + assert content[_MAX_CONTENT_INDEX]["text"] == "kept"