fix(guardrails): validate streamed tool context types

This commit is contained in:
albertbausili 2026-09-28 12:55:28 +02:00
parent 400f845b5f
commit 99fe53ab99

View file

@ -18,9 +18,11 @@ import json
import time
import uuid
from collections.abc import Mapping, Sequence
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Union, cast
from pydantic import TypeAdapter
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
@ -46,7 +48,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
unappliable_request_rewrite,
)
from litellm.main import stream_chunk_builder
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk, ChatCompletionToolParam
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
coerce_stream_holdback_value,
)
@ -796,12 +798,14 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
if self._streamed_tool_call_fingerprints(responses_so_far):
assembled: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
inputs["tool_calls"] = [
converted
for choice in assembled.choices
for tool_call in choice.message.tool_calls or ()
if (converted := self._convert_tool_call_to_dict(tool_call)) is not None
]
tool_calls: Final = chain.from_iterable(choice.message.tool_calls or () for choice in assembled.choices)
inputs["tool_calls"] = TypeAdapter(list[ChatCompletionToolCallChunk]).validate_python(
tuple(
{"index": index, **converted}
for index, tool_call in enumerate(tool_calls)
if (converted := self._convert_tool_call_to_dict(tool_call)) is not None
)
)
if responses_so_far and getattr(responses_so_far[0], "model", None):
inputs["model"] = responses_so_far[0].model
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(