mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(greptile): cap SSE content_index padding and use multiset tool-id check
This commit is contained in:
parent
3d0fc09f8a
commit
8f4c4630d3
4 changed files with 102 additions and 5 deletions
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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 -------------------------------------------
|
||||
|
||||
|
|
|
|||
57
tests/test_litellm/responses/test_sse_output_recovery.py
Normal file
57
tests/test_litellm/responses/test_sse_output_recovery.py
Normal file
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue