mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(guardrails): hold legacy function_call stream windows until the end-of-stream scan
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
0d605b7b45
commit
8acd2477a6
3 changed files with 104 additions and 26 deletions
|
|
@ -806,7 +806,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
stream_item_fingerprint(tool_call)
|
||||
for chunk in responses_so_far
|
||||
for choice in _stream_chunk_choices(chunk)
|
||||
for tool_call in stream_item_items(stream_item_field(choice, "delta"), "tool_calls")
|
||||
for tool_call in _streamed_delta_tool_calls(stream_item_field(choice, "delta"))
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1344,6 +1344,12 @@ def _stream_chunk_choices(item: object) -> Sequence[object]:
|
|||
return ()
|
||||
|
||||
|
||||
def _streamed_delta_tool_calls(delta: object) -> tuple[object, ...]:
|
||||
function_call: Final = stream_item_field(delta, "function_call")
|
||||
legacy: Final = () if function_call is None else (function_call,)
|
||||
return stream_item_items(delta, "tool_calls") + legacy
|
||||
|
||||
|
||||
def _blocked_stream_identity(
|
||||
exc: "ModifyResponseException", responses_so_far: Sequence[object]
|
||||
) -> tuple[str, int, str]:
|
||||
|
|
|
|||
|
|
@ -2213,6 +2213,28 @@ class TestStreamingScanKey:
|
|||
assert ended_key.tool_calls_in_flight is False
|
||||
assert ended_key != open_key
|
||||
|
||||
def test_legacy_function_call_delta_is_held_like_a_tool_call(self):
|
||||
from litellm.types.utils import Delta, FunctionCall, ModelResponseStream, StreamingChoices
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
function_chunk = ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=None, function_call=FunctionCall(name="run_shell", arguments='{"cmd": "rm"}')),
|
||||
finish_reason=None,
|
||||
)
|
||||
]
|
||||
)
|
||||
open_key = handler.get_streaming_scan_key([self._chunk("hi"), function_chunk])
|
||||
ended_key = handler.get_streaming_scan_key(
|
||||
[self._chunk("hi"), function_chunk, self._chunk(None, finish_reason="function_call")]
|
||||
)
|
||||
assert open_key.tool_calls_in_flight is True
|
||||
assert open_key.tool_calls == ()
|
||||
assert len(ended_key.tool_calls) == 1 and "run_shell" in ended_key.tool_calls[0]
|
||||
assert ended_key.tool_calls_in_flight is False
|
||||
|
||||
def test_text_after_the_first_choice_finishes_still_changes_the_key(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
first_done = [self._chunk("a", index=0), self._chunk("b", finish_reason="stop", index=0)]
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from litellm.types.utils import (
|
|||
ChatCompletionDeltaToolCall,
|
||||
Delta,
|
||||
Function,
|
||||
FunctionCall,
|
||||
GenericGuardrailAPIInputs,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
|
|
@ -219,29 +220,30 @@ def _chat_chunk(content: str = "", finish_reason: str | None = None) -> ModelRes
|
|||
)
|
||||
|
||||
|
||||
def _tool_call_chunk(arguments: str, finish_reason: str | None = None) -> ModelResponseStream:
|
||||
def _tool_call_chunk(
|
||||
arguments: str, finish_reason: str | None = None, legacy_function_call: bool = False
|
||||
) -> ModelResponseStream:
|
||||
delta = (
|
||||
Delta(role="assistant", content=None, function_call=FunctionCall(name="run_shell", arguments=arguments))
|
||||
if legacy_function_call
|
||||
else Delta(
|
||||
role="assistant",
|
||||
content=None,
|
||||
tool_calls=[
|
||||
ChatCompletionDeltaToolCall(
|
||||
id="call_1",
|
||||
type="function",
|
||||
index=0,
|
||||
function=Function(name="run_shell", arguments=arguments),
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-windowed",
|
||||
created=1724900000,
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(
|
||||
role="assistant",
|
||||
content=None,
|
||||
tool_calls=[
|
||||
ChatCompletionDeltaToolCall(
|
||||
id="call_1",
|
||||
type="function",
|
||||
index=0,
|
||||
function=Function(name="run_shell", arguments=arguments),
|
||||
)
|
||||
],
|
||||
),
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
],
|
||||
choices=[StreamingChoices(index=0, delta=delta, finish_reason=finish_reason)],
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -250,13 +252,14 @@ async def _windowed_chat_stream(
|
|||
collected: List[Any],
|
||||
content_chunks: List[str],
|
||||
tool_argument_chunks: List[str] | None = None,
|
||||
legacy_function_call: bool = False,
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
for content in content_chunks:
|
||||
yielded_count.append(len(collected))
|
||||
yield _chat_chunk(content)
|
||||
for arguments in tool_argument_chunks or []:
|
||||
yielded_count.append(len(collected))
|
||||
yield _tool_call_chunk(arguments)
|
||||
yield _tool_call_chunk(arguments, legacy_function_call=legacy_function_call)
|
||||
yielded_count.append(len(collected))
|
||||
yield _chat_chunk(finish_reason="tool_calls" if tool_argument_chunks else "stop")
|
||||
|
||||
|
|
@ -271,11 +274,22 @@ def _tool_argument_text(chunks: List[Any]) -> str:
|
|||
)
|
||||
|
||||
|
||||
def _function_call_argument_text(chunks: list[Any]) -> str:
|
||||
return "".join(
|
||||
choice.delta.function_call.arguments or ""
|
||||
for chunk in chunks
|
||||
if isinstance(chunk, ModelResponseStream)
|
||||
for choice in chunk.choices
|
||||
if choice.delta.function_call is not None
|
||||
)
|
||||
|
||||
|
||||
async def _run_windowed(
|
||||
guardrail: CustomGuardrail,
|
||||
content_chunks: List[str],
|
||||
end_of_stream_only: bool = False,
|
||||
tool_argument_chunks: List[str] | None = None,
|
||||
legacy_function_call: bool = False,
|
||||
) -> tuple[List[Any], List[int]]:
|
||||
guardrail.streaming_buffer_until_moderated = True
|
||||
guardrail.streaming_buffer_release_on_scan = True
|
||||
|
|
@ -292,7 +306,9 @@ async def _run_windowed(
|
|||
yielded_count: List[int] = []
|
||||
async for chunk in unified.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=_windowed_chat_stream(yielded_count, collected, content_chunks, tool_argument_chunks),
|
||||
response=_windowed_chat_stream(
|
||||
yielded_count, collected, content_chunks, tool_argument_chunks, legacy_function_call
|
||||
),
|
||||
request_data=request_data,
|
||||
):
|
||||
collected.append(chunk)
|
||||
|
|
@ -305,13 +321,24 @@ def _responses_message_stream_events(text_chunks: List[str]) -> List[dict]:
|
|||
return [
|
||||
{"type": "response.output_item.added", "output_index": 0, "item": {**message, "content": []}},
|
||||
*(
|
||||
{"type": "response.output_text.delta", "item_id": "msg_1", "output_index": 0, "content_index": 0, "delta": text}
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg_1",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": text,
|
||||
}
|
||||
for text in text_chunks
|
||||
),
|
||||
{"type": "response.output_item.done", "output_index": 0, "item": {**message, "content": content}},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {"id": "resp_1", "model": "gpt-4o", "status": "completed", "output": [{**message, "content": content}]},
|
||||
"response": {
|
||||
"id": "resp_1",
|
||||
"model": "gpt-4o",
|
||||
"status": "completed",
|
||||
"output": [{**message, "content": content}],
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
|
@ -322,7 +349,13 @@ def _responses_truncated_function_call_events(text: str, argument_chunks: List[s
|
|||
function_call = {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "run_shell"}
|
||||
return [
|
||||
{"type": "response.output_item.added", "output_index": 0, "item": {**message, "content": []}},
|
||||
{"type": "response.output_text.delta", "item_id": "msg_1", "output_index": 0, "content_index": 0, "delta": text},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg_1",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": text,
|
||||
},
|
||||
{"type": "response.output_item.added", "output_index": 1, "item": {**function_call, "arguments": ""}},
|
||||
*(
|
||||
{"type": "response.function_call_arguments.delta", "item_id": "fc_1", "output_index": 1, "delta": arguments}
|
||||
|
|
@ -463,6 +496,21 @@ async def test_windowed_buffer_holds_tool_call_windows_until_end_of_stream_scan(
|
|||
assert guardrail.tool_call_scan_indexes == [guardrail.scan_count]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_buffer_holds_legacy_function_call_windows_until_end_of_stream():
|
||||
guardrail = _PassingGuardrail(guardrail_name="windowed-functions", event_hook="post_call")
|
||||
content_chunks = ["one ", "two ", "three "]
|
||||
function_argument_chunks = ['{"cmd": "', TOOL_ARGUMENTS_MARKER, '"}']
|
||||
|
||||
collected, yielded_count = await _run_windowed(
|
||||
guardrail, content_chunks, tool_argument_chunks=function_argument_chunks, legacy_function_call=True
|
||||
)
|
||||
|
||||
assert yielded_count == [0, 0, 2, 2, 2, 2, 2]
|
||||
assert _chat_text(collected) == "".join(content_chunks)
|
||||
assert _function_call_argument_text(collected) == "".join(function_argument_chunks)
|
||||
|
||||
|
||||
def test_tool_call_only_scan_key_is_not_skipped_as_empty():
|
||||
assert _is_redundant_scan(StreamingScanKey(texts=("",)), None) is True
|
||||
assert _is_redundant_scan(StreamingScanKey(texts=("",), tool_calls=("run_shell:{}",)), None) is False
|
||||
|
|
@ -470,7 +518,9 @@ def test_tool_call_only_scan_key_is_not_skipped_as_empty():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_responses_output_item_done_round_keeps_text_window_withheld():
|
||||
guardrail = _MarkerBlockingGuardrail(guardrail_name="windowed-responses", event_hook="post_call", marker=ORIGINAL_MARKER)
|
||||
guardrail = _MarkerBlockingGuardrail(
|
||||
guardrail_name="windowed-responses", event_hook="post_call", marker=ORIGINAL_MARKER
|
||||
)
|
||||
events = _responses_message_stream_events(["one ", f"{ORIGINAL_MARKER} "])
|
||||
|
||||
raw = await _run_windowed_responses(guardrail, events)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue