mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): scan each choice's tool-call arguments apart on n>1 streams and log why a rewrite was discarded
The rebuilt streamed response keyed tool-call fragments by tool index alone, so on n>1 chat streams the two choices' argument fragments were concatenated into one string and post_call guardrails scanned garbled JSON. Fragments are now keyed by (choice index, tool index). When a guardrail's rewrite cannot be written back to the stream (multi-choice streams, a rewrite that adds or drops a tool call, legacy-hook shapes the translation cannot rescan), the pipeline now logs a warning naming the guardrail and the exact reason before releasing the original stream. Also commits the regenerated dashboard API types that make check produced.
This commit is contained in:
parent
30f33a949b
commit
20d80b5420
9 changed files with 254 additions and 66 deletions
|
|
@ -138,9 +138,13 @@ class _ToolCallDelta(TypedDict, total=False):
|
|||
|
||||
|
||||
class _ToolCallChoice(TypedDict, total=False):
|
||||
index: ReadOnly[int]
|
||||
delta: ReadOnly[_ToolCallDelta]
|
||||
|
||||
|
||||
_ToolCallKey: TypeAlias = tuple[int, int]
|
||||
|
||||
|
||||
class _ToolCallChunk(TypedDict):
|
||||
choices: ReadOnly[Sequence[_ToolCallChoice]]
|
||||
|
||||
|
|
@ -416,40 +420,41 @@ class ChunkProcessor:
|
|||
@staticmethod
|
||||
def _iter_tool_call_fragments(
|
||||
tool_call_chunks: Sequence["_ToolCallChunk"],
|
||||
) -> Iterator[tuple[int, str, str]]:
|
||||
) -> Iterator[tuple[_ToolCallKey, str, str]]:
|
||||
for chunk in tool_call_chunks:
|
||||
for choice in chunk["choices"]:
|
||||
delta = choice.get("delta")
|
||||
if not delta:
|
||||
continue
|
||||
choice_index = choice.get("index", 0)
|
||||
for tool_call in delta.get("tool_calls", ()):
|
||||
if not tool_call:
|
||||
continue
|
||||
if isinstance(tool_call, dict):
|
||||
index = tool_call.get("index", 0)
|
||||
key = (choice_index, tool_call.get("index", 0))
|
||||
function = tool_call.get("function")
|
||||
if isinstance(function, dict):
|
||||
if fragment_arguments := function.get("arguments"):
|
||||
yield index, "arguments", fragment_arguments
|
||||
yield key, "arguments", fragment_arguments
|
||||
elif function_arguments := getattr(function, "arguments", None):
|
||||
yield index, "arguments", function_arguments
|
||||
yield key, "arguments", function_arguments
|
||||
custom = tool_call.get("custom")
|
||||
if isinstance(custom, dict) and (custom_input := custom.get("input")):
|
||||
yield index, "custom_input", custom_input
|
||||
yield key, "custom_input", custom_input
|
||||
else:
|
||||
index = getattr(tool_call, "index", 0)
|
||||
key = (choice_index, getattr(tool_call, "index", 0))
|
||||
function = getattr(tool_call, "function", None)
|
||||
if object_arguments := getattr(function, "arguments", None):
|
||||
yield index, "arguments", object_arguments
|
||||
yield key, "arguments", object_arguments
|
||||
custom = getattr(tool_call, "custom", None)
|
||||
if object_custom_input := getattr(custom, "input", None):
|
||||
yield index, "custom_input", object_custom_input
|
||||
yield key, "custom_input", object_custom_input
|
||||
|
||||
@staticmethod
|
||||
def _join_fragments_by_index_and_field(
|
||||
fragment_records: Iterator[tuple[int, str, str]],
|
||||
) -> Mapping[tuple[int, str], str]:
|
||||
def group_key(record: tuple[int, str, str]) -> tuple[int, str]:
|
||||
def _join_fragments_by_key_and_field(
|
||||
fragment_records: Iterator[tuple[_ToolCallKey, str, str]],
|
||||
) -> Mapping[tuple[_ToolCallKey, str], str]:
|
||||
def group_key(record: tuple[_ToolCallKey, str, str]) -> tuple[_ToolCallKey, str]:
|
||||
return record[0], record[1]
|
||||
|
||||
return MappingProxyType(
|
||||
|
|
@ -467,13 +472,14 @@ class ChunkProcessor:
|
|||
tool_calls_list: list[
|
||||
ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall
|
||||
] = [] # mutable-ok: see return type
|
||||
tool_call_map: Final[dict[int, dict[str, Any]]] = {} # Map to store tool calls by index
|
||||
tool_call_map: Final[dict[_ToolCallKey, dict[str, Any]]] = {} # Map to store tool calls by choice and index
|
||||
|
||||
for chunk in tool_call_chunks:
|
||||
choices = chunk["choices"]
|
||||
for choice in choices:
|
||||
delta = choice.get("delta", {})
|
||||
tool_calls = delta.get("tool_calls", [])
|
||||
choice_index = choice.get("index", 0)
|
||||
|
||||
for tool_call in tool_calls:
|
||||
# Handle both dict and object formats
|
||||
|
|
@ -495,9 +501,9 @@ class ChunkProcessor:
|
|||
|
||||
# Get index (handle both dict and object)
|
||||
if isinstance(tool_call, dict):
|
||||
index = tool_call.get("index", 0)
|
||||
index = (choice_index, tool_call.get("index", 0))
|
||||
else:
|
||||
index = getattr(tool_call, "index", 0)
|
||||
index = (choice_index, getattr(tool_call, "index", 0))
|
||||
|
||||
if index not in tool_call_map:
|
||||
tool_call_map[index] = {
|
||||
|
|
@ -572,7 +578,7 @@ class ChunkProcessor:
|
|||
if isinstance(provider_fields, dict):
|
||||
merged_provider_fields.update(provider_fields)
|
||||
|
||||
joined_fragments: Final = self._join_fragments_by_index_and_field(
|
||||
joined_fragments: Final = self._join_fragments_by_key_and_field(
|
||||
self._iter_tool_call_fragments(tool_call_chunks)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1172,7 +1172,10 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if deliver_ended_stream_rewrites and unended_texts and tuple(unended_texts) != (string_so_far,):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name or "unknown")
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_to_apply.guardrail_name or "unknown",
|
||||
"the stream never reported a stop_reason, so the text rewrite has no assembled response to land on",
|
||||
)
|
||||
return responses_so_far
|
||||
|
||||
def _prepare_request_data(
|
||||
|
|
@ -1318,7 +1321,11 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if len(block_indices) != len(post_guardrail_tool_calls):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_name)
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name,
|
||||
f"the guardrail returned {len(post_guardrail_tool_calls)} tool calls for a stream that carried "
|
||||
f"{len(block_indices)} tool_use blocks",
|
||||
)
|
||||
rewrites_by_block: Final = MappingProxyType(
|
||||
{
|
||||
index: after
|
||||
|
|
|
|||
|
|
@ -1041,13 +1041,13 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
choice.index for response in responses_so_far for choice in response.choices
|
||||
)
|
||||
if len(stream_choice_indices) != 1:
|
||||
# stream_chunk_builder collapses every choice into one index-0
|
||||
# choice, so a rewrite of the rebuilt response cannot be attributed
|
||||
# back to a single choice on an n>1 stream: report it undeliverable
|
||||
# rather than deliver the rewrite on the wrong choice
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_name)
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name,
|
||||
f"the stream carries {len(stream_choice_indices)} choices and the rebuilt response's text rewrite "
|
||||
"cannot be attributed to one of them",
|
||||
)
|
||||
target_choice_index: Final = next(iter(stream_choice_indices))
|
||||
await self._apply_guardrail_responses_to_output_streaming(
|
||||
responses=responses_so_far,
|
||||
|
|
@ -1105,10 +1105,22 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
choice.index for response in responses_so_far for choice in response.choices
|
||||
)
|
||||
fragments_by_tool_call: Final = self._function_tool_call_fragments(responses_so_far)
|
||||
if len(stream_choice_indices) != 1 or len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
|
||||
if len(stream_choice_indices) != 1:
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_name)
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name,
|
||||
f"the stream carries {len(stream_choice_indices)} choices and tool-call rewrites are only written "
|
||||
"back on single-choice streams",
|
||||
)
|
||||
if len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name,
|
||||
f"the guardrail returned {len(post_guardrail_tool_calls)} tool calls for a stream that carried "
|
||||
f"{len(fragments_by_tool_call)}",
|
||||
)
|
||||
for before, (name, arguments), fragments in zip(
|
||||
pre_guardrail_tool_calls, post_guardrail_tool_calls, fragments_by_tool_call
|
||||
):
|
||||
|
|
|
|||
|
|
@ -957,7 +957,10 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
if deliver_ended_stream_rewrites and fallback_texts and tuple(fallback_texts) != (string_so_far,):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name or "unknown")
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_to_apply.guardrail_name or "unknown",
|
||||
"the stream carried no terminal response envelope to write the text rewrite back into",
|
||||
)
|
||||
return responses_so_far
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1070,7 +1073,11 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(guardrail_name)
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name,
|
||||
f"the guardrail returned {len(post_guardrail_tool_calls)} tool calls and the stream's "
|
||||
f"{len(tool_call_items)} function_call items could not be lined up with them by call_id",
|
||||
)
|
||||
for output_item, rewrite in (
|
||||
(output_item, rewrites_by_call_id[call_id])
|
||||
for output_item, call_id in zip(tool_call_items, call_ids)
|
||||
|
|
|
|||
|
|
@ -50,12 +50,13 @@ except ImportError:
|
|||
|
||||
|
||||
class UndeliverableStreamRewrite(Exception):
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
def __init__(self, guardrail_name: str, reason: str) -> None:
|
||||
super().__init__(
|
||||
f"Guardrail '{guardrail_name}' rewrote the streamed response in a way this endpoint's "
|
||||
"streaming pipeline cannot deliver"
|
||||
f"Guardrail '{guardrail_name}' rewrote the streamed response but the rewrite cannot be written "
|
||||
f"back to the stream: {reason}"
|
||||
)
|
||||
self.guardrail_name: Final = guardrail_name
|
||||
self.reason: Final = reason
|
||||
|
||||
|
||||
class UnappliableRequestRewrite(Exception):
|
||||
|
|
@ -91,8 +92,22 @@ def _rewrote(sent: tuple[object, ...] | None, returned: tuple[object, ...] | Non
|
|||
return sent is not None and returned is not None and returned != sent
|
||||
|
||||
|
||||
def _changed_count(sent: tuple[object, ...] | None, returned: tuple[object, ...] | None) -> bool:
|
||||
return sent is not None and returned is not None and len(returned) != len(sent)
|
||||
def _count_change(sent: tuple[object, ...] | None, returned: tuple[object, ...] | None) -> tuple[int, int] | None:
|
||||
if sent is None or returned is None or len(returned) == len(sent):
|
||||
return None
|
||||
return (len(sent), len(returned))
|
||||
|
||||
|
||||
def _tool_call_mismatch_reason(
|
||||
sent: tuple[tuple[object, object], ...] | None, returned: tuple[tuple[object, object], ...] | None
|
||||
) -> str | None:
|
||||
if sent == returned:
|
||||
return None
|
||||
sent_count: Final = len(sent or ())
|
||||
returned_count: Final = len(returned or ())
|
||||
if sent_count == returned_count:
|
||||
return "the legacy hook changed a tool call's name or arguments, which this path cannot write back"
|
||||
return f"the legacy hook returned {returned_count} tool calls for a stream that carried {sent_count}"
|
||||
|
||||
|
||||
_GuardrailMethodT = TypeVar("_GuardrailMethodT", bound=Callable[..., object])
|
||||
|
|
@ -119,7 +134,7 @@ class _StreamRewriteObserver(CustomGuardrail):
|
|||
self.inner: Final = inner
|
||||
self.rewrote_texts = False
|
||||
self.rewrote_tool_calls = False
|
||||
self.changed_tool_call_count = False
|
||||
self.tool_call_count_change: tuple[int, int] | None = None
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return self.inner.structured_messages_cover_full_request()
|
||||
|
|
@ -140,11 +155,22 @@ class _StreamRewriteObserver(CustomGuardrail):
|
|||
returned_tool_shapes: Final = _tool_call_shapes(outputs.get("tool_calls"))
|
||||
self.rewrote_texts = self.rewrote_texts or _rewrote(sent_texts, _text_snapshot(outputs.get("texts")))
|
||||
self.rewrote_tool_calls = self.rewrote_tool_calls or _rewrote(sent_tool_shapes, returned_tool_shapes)
|
||||
self.changed_tool_call_count = self.changed_tool_call_count or _changed_count(
|
||||
self.tool_call_count_change = self.tool_call_count_change or _count_change(
|
||||
sent_tool_shapes, returned_tool_shapes
|
||||
)
|
||||
return outputs
|
||||
|
||||
def discard_reason(self, deliver_rewrites: bool) -> str | None:
|
||||
if self.tool_call_count_change is not None:
|
||||
sent, returned = self.tool_call_count_change
|
||||
return (
|
||||
f"the guardrail returned {returned} tool calls for a stream that carried {sent}, and a rewrite "
|
||||
"that drops or adds a tool call cannot be written back"
|
||||
)
|
||||
if not deliver_rewrites and (self.rewrote_texts or self.rewrote_tool_calls):
|
||||
return "this endpoint's streaming pipeline does not write ended-stream rewrites back yet"
|
||||
return None
|
||||
|
||||
|
||||
class _ScannedTextRecorder(CustomGuardrail):
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
|
|
@ -209,13 +235,24 @@ class _LegacyHookStreamAdapter(CustomGuardrail):
|
|||
if rewrite is None:
|
||||
return inputs
|
||||
rescanned: Final = await self._rescan(rewrite, logging_obj)
|
||||
guardrail_name: Final = self.guardrail_name or "unknown"
|
||||
if rescanned is None:
|
||||
raise UndeliverableStreamRewrite(self.guardrail_name or "unknown")
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name, "the legacy hook's response could not be rescanned by this endpoint's translation"
|
||||
)
|
||||
rewritten: Final = rescanned.get("texts")
|
||||
if len(_scanned_texts(rewritten)) != len(_scanned_texts(inputs.get("texts"))):
|
||||
raise UndeliverableStreamRewrite(self.guardrail_name or "unknown")
|
||||
if _tool_call_shapes(rescanned.get("tool_calls")) != _tool_call_shapes(inputs.get("tool_calls")):
|
||||
raise UndeliverableStreamRewrite(self.guardrail_name or "unknown")
|
||||
returned_text_count: Final = len(_scanned_texts(rewritten))
|
||||
sent_text_count: Final = len(_scanned_texts(inputs.get("texts")))
|
||||
if returned_text_count != sent_text_count:
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name,
|
||||
f"the legacy hook returned {returned_text_count} texts for a stream that carried {sent_text_count}",
|
||||
)
|
||||
tool_call_mismatch: Final = _tool_call_mismatch_reason(
|
||||
_tool_call_shapes(inputs.get("tool_calls")), _tool_call_shapes(rescanned.get("tool_calls"))
|
||||
)
|
||||
if tool_call_mismatch is not None:
|
||||
raise UndeliverableStreamRewrite(guardrail_name, tool_call_mismatch)
|
||||
if not rewritten:
|
||||
return inputs
|
||||
rewritten_inputs: Final[GenericGuardrailAPIInputs] = {**inputs, "texts": rewritten}
|
||||
|
|
@ -262,14 +299,16 @@ def _prepare_hook_input(
|
|||
|
||||
def _release_original_chunks(
|
||||
guardrail_name: str,
|
||||
reason: str,
|
||||
streaming_chunks: list[object], # mutable-ok: shared buffered-stream chunks, restored in place
|
||||
originals: Sequence[object],
|
||||
) -> None:
|
||||
streaming_chunks[:] = originals # rebind-ok: the caller's buffer is the stream the client receives
|
||||
verbose_proxy_logger.warning(
|
||||
"Pipeline: guardrail '%s' rewrote the streamed response in a way this endpoint's streaming "
|
||||
"pipeline cannot deliver yet; the rewrite was discarded and the original stream released",
|
||||
"Pipeline: guardrail '%s' rewrote the streamed response but the rewrite could not be written back to "
|
||||
"the stream: %s. The whole rewrite, text rewrites included, was discarded and the original stream released",
|
||||
guardrail_name,
|
||||
reason,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -442,13 +481,12 @@ class PipelineExecutor:
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=hook_input,
|
||||
)
|
||||
except UndeliverableStreamRewrite:
|
||||
_release_original_chunks(step.guardrail, streaming_chunks, originals)
|
||||
except UndeliverableStreamRewrite as undeliverable:
|
||||
_release_original_chunks(step.guardrail, undeliverable.reason, streaming_chunks, originals)
|
||||
return
|
||||
if observer.changed_tool_call_count or (
|
||||
not deliver_rewrites and (observer.rewrote_texts or observer.rewrote_tool_calls)
|
||||
):
|
||||
_release_original_chunks(step.guardrail, streaming_chunks, originals)
|
||||
discard_reason: Final = observer.discard_reason(deliver_rewrites)
|
||||
if discard_reason is not None:
|
||||
_release_original_chunks(step.guardrail, discard_reason, streaming_chunks, originals)
|
||||
return
|
||||
if not callback.records_own_guardrail_information:
|
||||
add_guardrail_to_applied_guardrails_header(request_data=hook_input, guardrail_name=step.guardrail)
|
||||
|
|
|
|||
|
|
@ -1288,6 +1288,61 @@ def _tool_call_delta_chunk(tool_call: dict[str, object] | ChatCompletionDeltaToo
|
|||
return {"choices": [{"delta": {"tool_calls": [tool_call]}}]}
|
||||
|
||||
|
||||
def _choice_tool_call_delta_chunk(choice_index: int, tool_call: dict[str, object]) -> dict[str, object]:
|
||||
return {"choices": [{"index": choice_index, "delta": {"tool_calls": [tool_call]}}]}
|
||||
|
||||
|
||||
def test_get_combined_tool_content_keeps_each_choices_arguments_apart_when_choices_share_a_tool_index():
|
||||
processor = ChunkProcessor.__new__(ChunkProcessor)
|
||||
chunks = [
|
||||
_choice_tool_call_delta_chunk(0, {"index": 0, "id": "call_a", "type": "function", "function": {"name": "f"}}),
|
||||
_choice_tool_call_delta_chunk(1, {"index": 0, "id": "call_b", "type": "function", "function": {"name": "f"}}),
|
||||
_choice_tool_call_delta_chunk(0, {"index": 0, "function": {"arguments": '{"fruit": "pers'}}),
|
||||
_choice_tool_call_delta_chunk(1, {"index": 0, "function": {"arguments": '{"fruit": "dur'}}),
|
||||
_choice_tool_call_delta_chunk(0, {"index": 0, "function": {"arguments": 'immon"}'}}),
|
||||
_choice_tool_call_delta_chunk(1, {"index": 0, "function": {"arguments": 'ian"}'}}),
|
||||
]
|
||||
|
||||
combined = processor.get_combined_tool_content(chunks)
|
||||
|
||||
assert [(tool_call.id, tool_call.function.arguments) for tool_call in combined] == [
|
||||
("call_a", '{"fruit": "persimmon"}'),
|
||||
("call_b", '{"fruit": "durian"}'),
|
||||
]
|
||||
|
||||
|
||||
def test_stream_chunk_builder_keeps_each_choices_tool_call_arguments_apart():
|
||||
def chunk(choice_index: int, tool_call: ChatCompletionDeltaToolCall) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-123",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567890,
|
||||
model="gpt-4.1-mini",
|
||||
choices=[StreamingChoices(index=choice_index, delta=Delta(tool_calls=[tool_call]), finish_reason=None)],
|
||||
)
|
||||
|
||||
def fragment(arguments: str, name: str | None = None, call_id: str | None = None) -> ChatCompletionDeltaToolCall:
|
||||
return ChatCompletionDeltaToolCall(
|
||||
id=call_id, index=0, type="function", function=Function(name=name, arguments=arguments)
|
||||
)
|
||||
|
||||
response = stream_chunk_builder(
|
||||
chunks=[
|
||||
chunk(0, fragment("", name="lookup_fruit", call_id="call_a")),
|
||||
chunk(1, fragment("", name="lookup_fruit", call_id="call_b")),
|
||||
chunk(0, fragment('{"fruit": "pers')),
|
||||
chunk(1, fragment('{"fruit": "dur')),
|
||||
chunk(0, fragment('immon"}')),
|
||||
chunk(1, fragment('ian"}')),
|
||||
]
|
||||
)
|
||||
|
||||
assert [(tool_call.id, tool_call.function.arguments) for tool_call in response.choices[0].message.tool_calls] == [
|
||||
("call_a", '{"fruit": "persimmon"}'),
|
||||
("call_b", '{"fruit": "durian"}'),
|
||||
]
|
||||
|
||||
|
||||
def test_get_combined_tool_content_joins_many_dict_shaped_argument_fragments_in_order():
|
||||
processor = ChunkProcessor.__new__(ChunkProcessor)
|
||||
first_fragments = [f"a{i};" for i in range(300)]
|
||||
|
|
|
|||
|
|
@ -1267,7 +1267,7 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = self._two_choice_stream_chunks()
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite):
|
||||
with pytest.raises(UndeliverableStreamRewrite, match="the stream carries 2 choices") as raised:
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=self._world_masking_guardrail(),
|
||||
|
|
@ -1275,6 +1275,11 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert raised.value.guardrail_name == "test-mask"
|
||||
assert raised.value.reason == (
|
||||
"the stream carries 2 choices and the rebuilt response's text rewrite cannot be attributed to one of them"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _two_choice_tool_call_stream_chunks() -> list:
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -1310,12 +1315,51 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
return [
|
||||
chunk(0, fragment("", name="lookup_fruit", call_id="call_1")),
|
||||
chunk(1, fragment("", name="lookup_fruit", call_id="call_2")),
|
||||
chunk(0, fragment('{"fruit": "persimmon"}')),
|
||||
chunk(1, fragment('{"fruit": "durian"}')),
|
||||
chunk(0, fragment('{"fruit": "pers')),
|
||||
chunk(1, fragment('{"fruit": "dur')),
|
||||
chunk(0, fragment('immon"}')),
|
||||
chunk(1, fragment('ian"}')),
|
||||
chunk(0, None, finish_reason="tool_calls"),
|
||||
chunk(1, None, finish_reason="tool_calls"),
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _recording_guardrail() -> CustomGuardrail:
|
||||
class Recorder(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(guardrail_name="recorder")
|
||||
self.seen_inputs: list[GenericGuardrailAPIInputs] = []
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.seen_inputs.append(inputs)
|
||||
return inputs
|
||||
|
||||
return Recorder()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ended_multi_choice_stream_scans_each_choices_tool_call_arguments_apart(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = self._two_choice_tool_call_stream_chunks()
|
||||
guardrail = self._recording_guardrail()
|
||||
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=guardrail,
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert [
|
||||
(tool_call["id"], tool_call["function"]["arguments"])
|
||||
for tool_call in guardrail.seen_inputs[-1]["tool_calls"]
|
||||
] == [("call_1", '{"fruit": "persimmon"}'), ("call_2", '{"fruit": "durian"}')]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_ended_stream_tool_call_rewrite_on_multi_choice_stream_fails_closed(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
|
@ -1323,7 +1367,7 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = self._two_choice_tool_call_stream_chunks()
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite):
|
||||
with pytest.raises(UndeliverableStreamRewrite, match="the stream carries 2 choices") as raised:
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=MockGuardrail(guardrail_name="test"),
|
||||
|
|
@ -1331,6 +1375,11 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert raised.value.guardrail_name == "test"
|
||||
assert raised.value.reason == (
|
||||
"the stream carries 2 choices and tool-call rewrites are only written back on single-choice streams"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_ended_stream_clean_multi_choice_stream_released_untouched(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
|
|
|
|||
|
|
@ -1122,7 +1122,7 @@ class _RefusingTranslation:
|
|||
deliver_ended_stream_rewrites=False,
|
||||
):
|
||||
responses_so_far[0]["text"] = "half-written"
|
||||
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name)
|
||||
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name, "the translation refused it")
|
||||
|
||||
|
||||
def _chunk():
|
||||
|
|
@ -1143,10 +1143,22 @@ async def _run_streaming_step(translation, streaming_chunks=None):
|
|||
)
|
||||
|
||||
|
||||
def _assert_passed_with_discard_warning(result, caplog):
|
||||
NO_WRITE_BACK_REASON = "this endpoint's streaming pipeline does not write ended-stream rewrites back yet"
|
||||
|
||||
|
||||
def _assert_passed_with_discard_warning(result, caplog, reason):
|
||||
assert result.terminal_action == "allow"
|
||||
assert [step.outcome for step in result.step_results] == ["pass"]
|
||||
assert any("'masker'" in record.getMessage() and "discarded" in record.getMessage() for record in caplog.records)
|
||||
discard_warnings = [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if record.levelno == logging.WARNING
|
||||
and "'masker'" in record.getMessage()
|
||||
and "discarded" in record.getMessage()
|
||||
]
|
||||
assert len(discard_warnings) == 1
|
||||
assert reason in discard_warnings[0]
|
||||
assert "text rewrites included" in discard_warnings[0]
|
||||
assert "masker" not in ((result.modified_data or {}).get("metadata") or {}).get("applied_guardrails", [])
|
||||
|
||||
|
||||
|
|
@ -1159,7 +1171,7 @@ async def test_streaming_step_discards_text_rewrite_when_translation_lacks_write
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_streaming_step(translation, chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, NO_WRITE_BACK_REASON)
|
||||
assert chunks == [_chunk()]
|
||||
assert translation.seen_guardrail_names == ["masker"]
|
||||
|
||||
|
|
@ -1196,7 +1208,7 @@ async def test_streaming_step_in_place_rewrite_is_discarded_without_write_back(m
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_streaming_step(_TextTranslation(), chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, NO_WRITE_BACK_REASON)
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1258,7 +1270,9 @@ async def test_streaming_step_discards_whole_rewrite_when_guardrail_drops_a_tool
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_streaming_step(_WritingTranslation(), chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(
|
||||
result, caplog, "the guardrail returned 0 tool calls for a stream that carried 1"
|
||||
)
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1270,7 +1284,7 @@ async def test_streaming_step_discards_tool_call_rewrite_when_translation_lacks_
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_streaming_step(_TextTranslation(), chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, NO_WRITE_BACK_REASON)
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1326,7 +1340,7 @@ async def test_streaming_step_restores_chunks_when_translation_refuses_the_rewri
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_streaming_step(_RefusingTranslation(), chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, "the translation refused it")
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1561,7 +1575,7 @@ async def test_streaming_step_discards_legacy_rewrite_whose_texts_do_not_line_up
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, "the legacy hook returned 2 texts for a stream that carried 1")
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1576,7 +1590,7 @@ async def test_streaming_step_discards_legacy_rewrite_that_changes_a_tool_call(m
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, "the legacy hook changed a tool call's name or arguments")
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1588,7 +1602,9 @@ async def test_streaming_step_discards_legacy_rewrite_that_drops_the_tool_calls(
|
|||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(
|
||||
result, caplog, "the legacy hook returned 0 tool calls for a stream that carried 1"
|
||||
)
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1620,7 +1636,7 @@ async def test_streaming_step_discards_a_legacy_tool_call_rewrite_on_a_tool_only
|
|||
monkeypatch, guardrail, chunks, translation=_ToolOnlyLegacyScanningTranslation()
|
||||
)
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, "the legacy hook changed a tool call's name or arguments")
|
||||
assert chunks == [_tool_only_chunk()]
|
||||
|
||||
|
||||
|
|
@ -1658,7 +1674,7 @@ async def test_streaming_step_discards_a_legacy_rewrite_the_translation_cannot_r
|
|||
|
||||
result = await _run_legacy_streaming_step(monkeypatch, guardrail, chunks, translation=_UnscannableRewriteTranslation())
|
||||
|
||||
_assert_passed_with_discard_warning(result, caplog)
|
||||
_assert_passed_with_discard_warning(result, caplog, "the legacy hook's response could not be rescanned")
|
||||
assert chunks == [_chunk()]
|
||||
|
||||
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -16781,7 +16781,6 @@ export interface paths {
|
|||
* - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking.
|
||||
* - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" }
|
||||
* - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x.
|
||||
* - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests.
|
||||
* - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys)
|
||||
* - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
|
||||
* - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
|
||||
|
|
@ -16887,7 +16886,6 @@ export interface paths {
|
|||
* - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking.
|
||||
* - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" }
|
||||
* - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x.
|
||||
* - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests.
|
||||
* - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys)
|
||||
* - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
|
||||
* - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue