mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): validate streamed tool context types
This commit is contained in:
parent
400f845b5f
commit
99fe53ab99
1 changed files with 11 additions and 7 deletions
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue