mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge 2354bf4cad into dab2deb5ed
This commit is contained in:
commit
0790ed731c
6 changed files with 594 additions and 17 deletions
|
|
@ -114,7 +114,12 @@ from litellm.proxy.dd_span_tagger import DDSpanTagger
|
|||
from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression
|
||||
from litellm.proxy.native_compaction import with_proxy_compaction_executor
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
|
||||
from litellm.proxy.utils import (
|
||||
ProxyLogging,
|
||||
StreamingToolCallState,
|
||||
_check_and_merge_model_level_guardrails,
|
||||
streaming_tool_calls_with_response,
|
||||
)
|
||||
from litellm.router import Router
|
||||
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
|
||||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
|
|
@ -3856,6 +3861,27 @@ class ProxyBaseLLMRequestProcessing:
|
|||
):
|
||||
await logging_obj.invalidate_baseline_cache_estimate("incomplete_response", completed=True)
|
||||
|
||||
@staticmethod
|
||||
async def _apply_streaming_chunk_hook(
|
||||
*,
|
||||
chunk: Any, # noqa: ANN401 # streaming callbacks may replace chunks with provider-specific response types
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_data: dict, # mutable-ok: ProxyLogging requires the existing mutable request payload
|
||||
str_so_far: str,
|
||||
streaming_tool_calls_so_far: StreamingToolCallState,
|
||||
) -> tuple[Any, StreamingToolCallState]:
|
||||
return (
|
||||
await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=chunk,
|
||||
data=request_data,
|
||||
str_so_far=str_so_far,
|
||||
streaming_tool_calls_so_far=streaming_tool_calls_so_far,
|
||||
),
|
||||
streaming_tool_calls_with_response(streaming_tool_calls_so_far, chunk),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def async_streaming_data_generator(
|
||||
response: object,
|
||||
|
|
@ -3902,6 +3928,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes
|
||||
try:
|
||||
str_so_far = ""
|
||||
streaming_tool_calls_so_far: StreamingToolCallState = ()
|
||||
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=response,
|
||||
|
|
@ -3913,11 +3940,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
verbose_proxy_logger.debug("async_data_generator: received streaming chunk - %s", chunk)
|
||||
|
||||
if not fast_path:
|
||||
chunk = await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
(
|
||||
chunk,
|
||||
streaming_tool_calls_so_far,
|
||||
) = await ProxyBaseLLMRequestProcessing._apply_streaming_chunk_hook(
|
||||
chunk=chunk,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=chunk,
|
||||
data=request_data,
|
||||
request_data=request_data,
|
||||
str_so_far=str_so_far,
|
||||
streaming_tool_calls_so_far=streaming_tool_calls_so_far,
|
||||
)
|
||||
|
||||
if isinstance(chunk, (ModelResponse, ModelResponseStream)):
|
||||
|
|
|
|||
|
|
@ -793,6 +793,7 @@ from litellm.proxy.utils import (
|
|||
PrismaClient,
|
||||
ProxyLogging,
|
||||
ProxyUpdateSpend,
|
||||
StreamingToolCallState,
|
||||
_cache_user_row,
|
||||
_get_docs_url,
|
||||
_get_openapi_url,
|
||||
|
|
@ -813,6 +814,7 @@ from litellm.proxy.utils import (
|
|||
migrate_passwords_to_scrypt_async,
|
||||
model_dump_with_preserved_fields,
|
||||
prefetch_config_params,
|
||||
streaming_tool_calls_with_response,
|
||||
update_spend,
|
||||
)
|
||||
from litellm.proxy.video_endpoints.endpoints import router as video_router
|
||||
|
|
@ -9338,19 +9340,23 @@ async def _apply_streaming_chunk_hooks(
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_data: dict,
|
||||
str_so_far: str,
|
||||
) -> tuple[Any, str]:
|
||||
streaming_tool_calls_so_far: StreamingToolCallState = (),
|
||||
) -> tuple[Any, str, StreamingToolCallState]:
|
||||
stream_chunk: Final = chunk
|
||||
chunk = await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=chunk,
|
||||
data=request_data,
|
||||
str_so_far=str_so_far if str_so_far else None,
|
||||
streaming_tool_calls_so_far=streaming_tool_calls_so_far,
|
||||
)
|
||||
|
||||
if isinstance(chunk, (ModelResponse, ModelResponseStream)):
|
||||
response_str: Final = litellm.get_response_string(response_obj=chunk)
|
||||
str_so_far += response_str
|
||||
|
||||
return chunk, str_so_far
|
||||
updated_tool_calls: Final = streaming_tool_calls_with_response(streaming_tool_calls_so_far, stream_chunk)
|
||||
return chunk, str_so_far, updated_tool_calls
|
||||
|
||||
|
||||
def _format_streaming_sse_chunk(chunk: str | bytes) -> str | bytes:
|
||||
|
|
@ -9607,6 +9613,7 @@ async def async_data_generator(
|
|||
# Previously "".join(str_so_far_parts) was called every chunk, re-joining
|
||||
# the entire accumulated response. String += is O(n) amortized total.
|
||||
_str_so_far: str = ""
|
||||
_streaming_tool_calls_so_far: StreamingToolCallState = ()
|
||||
# Separate iterator-level vs per-chunk hook decisions. The iterator
|
||||
# wrap is needed when any callback overrides
|
||||
# ``async_post_call_streaming_iterator_hook`` or has
|
||||
|
|
@ -9656,11 +9663,12 @@ async def async_data_generator(
|
|||
chunk = cast(Any, item) # cast-ok: sentinel already handled above, item is a real chunk here
|
||||
if needs_per_chunk_hook:
|
||||
### CALL HOOKS ### - modify outgoing data
|
||||
chunk, _str_so_far = await _apply_streaming_chunk_hooks(
|
||||
chunk, _str_so_far, _streaming_tool_calls_so_far = await _apply_streaming_chunk_hooks(
|
||||
chunk=chunk,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
str_so_far=_str_so_far,
|
||||
streaming_tool_calls_so_far=_streaming_tool_calls_so_far,
|
||||
)
|
||||
|
||||
# Mid-stream fallbacks surface metadata on individual chunks rather than
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from collections.abc import (
|
|||
Awaitable,
|
||||
Callable,
|
||||
Coroutine,
|
||||
Iterator,
|
||||
Mapping,
|
||||
Sequence,
|
||||
)
|
||||
|
|
@ -246,7 +247,12 @@ from litellm.types.mcp import (
|
|||
MCPPreCallResponseObject,
|
||||
)
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult
|
||||
from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionDeltaCustomToolCall,
|
||||
ChatCompletionDeltaToolCall,
|
||||
LLMResponseTypes,
|
||||
LoggedLiteLLMParams,
|
||||
)
|
||||
from litellm.utils import (
|
||||
_add_custom_logger_callback_to_specific_event, # pyright: ignore[reportPrivateUsage] # only string-to-logger helper
|
||||
)
|
||||
|
|
@ -1084,6 +1090,127 @@ def _call_type_for_route(route: str | None) -> str | None:
|
|||
return call_types[0].value if len(operations) == 1 else None
|
||||
|
||||
|
||||
class _StreamingHookResponseText(str):
|
||||
"""Marks the exact text object passed to a per-chunk streaming hook."""
|
||||
|
||||
|
||||
class _StructuredStreamingGuardrailText(_StreamingHookResponseText):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _StreamingToolCallFragment:
|
||||
choice_index: int
|
||||
tool_index: int
|
||||
name: str
|
||||
arguments: str
|
||||
|
||||
@property
|
||||
def key(self) -> tuple[int, int]:
|
||||
return self.choice_index, self.tool_index
|
||||
|
||||
|
||||
StreamingToolCallState: TypeAlias = tuple[_StreamingToolCallFragment, ...]
|
||||
|
||||
|
||||
def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, response: object) -> str:
|
||||
complete_response: Final = str_so_far + response_str if str_so_far is not None else response_str
|
||||
if complete_response == "" and isinstance(response, (ModelResponse, ModelResponseStream)):
|
||||
return _StreamingHookResponseText(complete_response)
|
||||
return complete_response
|
||||
|
||||
|
||||
def _streaming_tool_call_fragments(response: ModelResponseStream) -> Iterator[_StreamingToolCallFragment]:
|
||||
for choice in response.choices:
|
||||
for tool_call in choice.delta.tool_calls or ():
|
||||
if isinstance(tool_call, ChatCompletionDeltaToolCall):
|
||||
yield _StreamingToolCallFragment(
|
||||
choice_index=choice.index,
|
||||
tool_index=tool_call.index,
|
||||
name=tool_call.function.name or "",
|
||||
arguments=tool_call.function.arguments or "",
|
||||
)
|
||||
elif isinstance(tool_call, ChatCompletionDeltaCustomToolCall):
|
||||
yield _StreamingToolCallFragment(
|
||||
choice_index=choice.index,
|
||||
tool_index=tool_call.index,
|
||||
name=tool_call.custom.name or "",
|
||||
arguments=tool_call.custom.input or "",
|
||||
)
|
||||
if choice.delta.function_call is not None:
|
||||
yield _StreamingToolCallFragment(
|
||||
choice_index=choice.index,
|
||||
tool_index=-1,
|
||||
name=choice.delta.function_call.name or "",
|
||||
arguments=choice.delta.function_call.arguments or "",
|
||||
)
|
||||
|
||||
|
||||
def _assembled_streaming_tool_calls(
|
||||
fragments: Sequence[_StreamingToolCallFragment],
|
||||
) -> StreamingToolCallState:
|
||||
keys: Final = tuple(
|
||||
fragment.key
|
||||
for position, fragment in enumerate(fragments)
|
||||
if fragment.key not in tuple(previous.key for previous in fragments[:position])
|
||||
)
|
||||
return tuple(
|
||||
_StreamingToolCallFragment(
|
||||
choice_index=key[0],
|
||||
tool_index=key[1],
|
||||
name="".join(fragment.name for fragment in fragments if fragment.key == key),
|
||||
arguments="".join(fragment.arguments for fragment in fragments if fragment.key == key),
|
||||
)
|
||||
for key in keys
|
||||
)
|
||||
|
||||
|
||||
def streaming_tool_calls_with_response(
|
||||
tool_calls: Sequence[_StreamingToolCallFragment], response: object
|
||||
) -> StreamingToolCallState:
|
||||
if not isinstance(response, ModelResponseStream):
|
||||
return tuple(tool_calls)
|
||||
fragments: Final = _streaming_tool_call_fragments(response)
|
||||
return _assembled_streaming_tool_calls((*tool_calls, *fragments))
|
||||
|
||||
|
||||
def _streaming_guardrail_response_text(
|
||||
*,
|
||||
complete_response: str,
|
||||
response: object,
|
||||
streaming_tool_calls_so_far: Sequence[_StreamingToolCallFragment],
|
||||
) -> str:
|
||||
if not isinstance(response, ModelResponseStream):
|
||||
return complete_response
|
||||
tool_calls: Final = streaming_tool_calls_with_response(streaming_tool_calls_so_far, response)
|
||||
structured_fields: Final = tuple(
|
||||
"tool_call:"
|
||||
+ json.dumps(
|
||||
(tool_call.choice_index, tool_call.tool_index, tool_call.name, tool_call.arguments),
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
for tool_call in tool_calls
|
||||
)
|
||||
if not structured_fields:
|
||||
return complete_response
|
||||
return _StructuredStreamingGuardrailText("\n".join((str(complete_response), *structured_fields)))
|
||||
|
||||
|
||||
def _is_unchanged_structured_streaming_hook_response(
|
||||
*, callback_response: object, complete_response: str, response_str: str, response: object
|
||||
) -> bool:
|
||||
if not isinstance(response, (ModelResponse, ModelResponseStream)):
|
||||
return False
|
||||
if isinstance(complete_response, _StructuredStreamingGuardrailText):
|
||||
return callback_response == complete_response
|
||||
if isinstance(complete_response, _StreamingHookResponseText):
|
||||
return callback_response is complete_response
|
||||
if response_str != "":
|
||||
return False
|
||||
return callback_response == complete_response
|
||||
|
||||
|
||||
_PROXY_ONLY_LLM_API_ERRORS: Final = (HTTPException, ProxyException, GuardrailRaisedException)
|
||||
|
||||
|
||||
|
|
@ -3806,6 +3933,7 @@ class ProxyLogging:
|
|||
response: ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
str_so_far: str | None = None,
|
||||
streaming_tool_calls_so_far: Sequence[_StreamingToolCallFragment] = (),
|
||||
):
|
||||
"""
|
||||
Allow user to modify outgoing streaming data -> per chunk
|
||||
|
|
@ -3872,18 +4000,30 @@ class ProxyLogging:
|
|||
else:
|
||||
_callback = callback
|
||||
if _callback is not None and isinstance(_callback, CustomLogger):
|
||||
if str_so_far is not None:
|
||||
complete_response = str_so_far + response_str
|
||||
else:
|
||||
complete_response = response_str
|
||||
complete_response = _streaming_hook_response_text(
|
||||
response_str=response_str,
|
||||
str_so_far=str_so_far,
|
||||
response=response,
|
||||
)
|
||||
if isinstance(_callback, CustomGuardrail):
|
||||
complete_response = _streaming_guardrail_response_text(
|
||||
complete_response=complete_response,
|
||||
response=response,
|
||||
streaming_tool_calls_so_far=streaming_tool_calls_so_far,
|
||||
)
|
||||
callback_response: (
|
||||
ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None
|
||||
str | ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None
|
||||
)
|
||||
callback_response = await _callback.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=complete_response,
|
||||
)
|
||||
if callback_response is not None:
|
||||
if callback_response is not None and not _is_unchanged_structured_streaming_hook_response(
|
||||
callback_response=callback_response,
|
||||
complete_response=complete_response,
|
||||
response_str=response_str,
|
||||
response=response,
|
||||
):
|
||||
response = callback_response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ from pydantic import BaseModel
|
|||
import litellm
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_streaming_chunk_hooks,
|
||||
|
|
@ -78,6 +79,30 @@ def _simple_chunk(model: str = "gpt-4", content: str = "hi") -> ModelResponseStr
|
|||
)
|
||||
|
||||
|
||||
def _tool_call_chunk(arguments: str) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-tool-call",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"type": "function",
|
||||
"function": {"name": "run_command", "arguments": arguments},
|
||||
}
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
created=0,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
|
||||
async def _async_iter(items):
|
||||
for it in items:
|
||||
yield it
|
||||
|
|
@ -567,12 +592,12 @@ def test_serialize_streaming_chunk_invalid_input_raises_attribute_error():
|
|||
async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch):
|
||||
chunk = _simple_chunk(content="abc")
|
||||
|
||||
async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None):
|
||||
async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None, streaming_tool_calls_so_far=()):
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough)
|
||||
|
||||
new_chunk, new_str = await _apply_streaming_chunk_hooks(
|
||||
new_chunk, new_str, tool_calls = await _apply_streaming_chunk_hooks(
|
||||
chunk=chunk,
|
||||
user_api_key_dict=_user_auth(),
|
||||
request_data={},
|
||||
|
|
@ -583,14 +608,37 @@ async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch):
|
|||
"chunk_is_basemodel": isinstance(new_chunk, ModelResponseStream),
|
||||
"str_so_far": new_str,
|
||||
"grew": len(new_str) > len("prior:"),
|
||||
"tool_calls": tool_calls,
|
||||
}
|
||||
assert observed == {
|
||||
"chunk_is_basemodel": True,
|
||||
"str_so_far": "prior:abc",
|
||||
"grew": True,
|
||||
"tool_calls": (),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_streaming_chunk_hooks_compacts_tool_call_history(monkeypatch):
|
||||
async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None, streaming_tool_calls_so_far=()):
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough)
|
||||
stream_state = ()
|
||||
|
||||
for arguments in ('{"command":"', "blocked-", "command", '"}'):
|
||||
_, _, stream_state = await _apply_streaming_chunk_hooks(
|
||||
chunk=_tool_call_chunk(arguments),
|
||||
user_api_key_dict=_user_auth(),
|
||||
request_data={},
|
||||
str_so_far="",
|
||||
streaming_tool_calls_so_far=stream_state,
|
||||
)
|
||||
|
||||
assert len(stream_state) == 1
|
||||
assert stream_state[0].arguments == '{"command":"blocked-command"}'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_streaming_chunk_hooks_hook_raises_exception(monkeypatch):
|
||||
async def _boom(*args, **kwargs):
|
||||
|
|
@ -693,6 +741,80 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch):
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_data_generator_preserves_tool_calls_through_per_chunk_hook(
|
||||
monkeypatch,
|
||||
):
|
||||
class _PassThrough(CustomLogger):
|
||||
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
|
||||
return response
|
||||
|
||||
callback = _PassThrough()
|
||||
monkeypatch.setattr(litellm, "callbacks", [callback])
|
||||
_patch_logging_flags(monkeypatch, needs_per_chunk=True)
|
||||
|
||||
chunk = ModelResponseStream(
|
||||
id="chatcmpl-tools",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call-weather",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": "",
|
||||
},
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"id": "call-time",
|
||||
"type": "function",
|
||||
"function": {"name": "get_time", "arguments": ""},
|
||||
},
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
created=0,
|
||||
model="gpt-4o-mini",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
out = []
|
||||
async for line in async_data_generator(
|
||||
response=_async_iter([chunk]),
|
||||
user_api_key_dict=_user_auth(),
|
||||
request_data={"model": "gpt-4o-mini"},
|
||||
):
|
||||
out.append(line)
|
||||
|
||||
first = out[0]
|
||||
assert isinstance(first, (str, bytes))
|
||||
first_text = first.decode() if isinstance(first, bytes) else first
|
||||
assert first_text.startswith("data: {")
|
||||
payload = json.loads(first_text.removeprefix("data: ").removesuffix("\n\n"))
|
||||
assert payload["choices"][0]["delta"]["tool_calls"] == [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call-weather",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": ""},
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"id": "call-time",
|
||||
"type": "function",
|
||||
"function": {"name": "get_time", "arguments": ""},
|
||||
},
|
||||
]
|
||||
assert out[-1] == "data: [DONE]\n\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_data_generator_uses_response_fallback_metadata(monkeypatch):
|
||||
_patch_logging_flags(monkeypatch)
|
||||
|
|
|
|||
|
|
@ -4756,6 +4756,48 @@ class TestAsyncStreamingDataGeneratorFastPath:
|
|||
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_streaming_chunk_hook_compacts_tool_call_history(self):
|
||||
proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||||
proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["response"])
|
||||
|
||||
def _chunk(arguments: str):
|
||||
return litellm.ModelResponseStream(
|
||||
id="chatcmpl-tool-call",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"type": "function",
|
||||
"function": {"name": "run_command", "arguments": arguments},
|
||||
}
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
created=0,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
stream_state = ()
|
||||
for arguments in ("blocked-", "command"):
|
||||
_, stream_state = await ProxyBaseLLMRequestProcessing._apply_streaming_chunk_hook(
|
||||
chunk=_chunk(arguments),
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=MagicMock(spec=ProxyUserAPIKeyAuth),
|
||||
request_data={},
|
||||
str_so_far="",
|
||||
streaming_tool_calls_so_far=stream_state,
|
||||
)
|
||||
|
||||
assert len(stream_state) == 1
|
||||
assert stream_state[0].arguments == "blocked-command"
|
||||
|
||||
|
||||
class TestDisconnectGatherCleanup:
|
||||
def _disconnect_request(self) -> Request:
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import (
|
|||
BaseAnthropicMessagesStreamingIterator,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.proxy.utils import ProxyLogging, streaming_tool_calls_with_response
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
|
@ -384,6 +384,239 @@ async def test_async_post_call_streaming_hook_invokes_per_chunk_callback(proxy_l
|
|||
assert out.startswith("modified-")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("str_so_far", [None, "I will check that. "])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_streaming_hook_preserves_tool_calls_when_callback_returns_unmodified_text(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch, str_so_far
|
||||
):
|
||||
class _PassThrough(CustomLogger):
|
||||
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PassThrough()])
|
||||
|
||||
response = litellm.ModelResponseStream(
|
||||
id="chatcmpl-tools",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call-weather",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": ""},
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"id": "call-time",
|
||||
"type": "function",
|
||||
"function": {"name": "get_time", "arguments": ""},
|
||||
},
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
created=0,
|
||||
model="gpt-4o-mini",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
out = await proxy_logging.async_post_call_streaming_hook(
|
||||
data={},
|
||||
response=response,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
|
||||
assert out is response
|
||||
assert [
|
||||
tool_call.model_dump(exclude_none=True)
|
||||
for tool_call in out.choices[0].delta.tool_calls
|
||||
] == [
|
||||
{
|
||||
"id": "call-weather",
|
||||
"function": {"arguments": "", "name": "get_weather"},
|
||||
"type": "function",
|
||||
"index": 0,
|
||||
},
|
||||
{
|
||||
"id": "call-time",
|
||||
"function": {"arguments": "", "name": "get_time"},
|
||||
"type": "function",
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
|
||||
class _Replace(CustomLogger):
|
||||
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
|
||||
return "replacement"
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PassThrough(), _Replace()])
|
||||
replaced = await proxy_logging.async_post_call_streaming_hook(
|
||||
data={},
|
||||
response=response,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
assert replaced == "replacement"
|
||||
|
||||
if str_so_far is not None:
|
||||
class _EquivalentCopy(CustomLogger):
|
||||
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
|
||||
return response.encode().decode()
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_EquivalentCopy()])
|
||||
equivalent = await proxy_logging.async_post_call_streaming_hook(
|
||||
data={},
|
||||
response=response,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
assert equivalent is response
|
||||
|
||||
class _Suppress(CustomLogger):
|
||||
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
|
||||
return ""
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_Suppress()])
|
||||
suppressed = await proxy_logging.async_post_call_streaming_hook(
|
||||
data={},
|
||||
response=response,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
assert suppressed == ""
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_streaming_hook_exposes_tool_calls_to_guardrails(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
class _BlockingGuardrail(CustomGuardrail):
|
||||
def should_run_guardrail(self, data, event_type):
|
||||
return True
|
||||
|
||||
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
|
||||
if "blocked-command" in response:
|
||||
return "data: blocked\n\n"
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"callbacks",
|
||||
[_BlockingGuardrail(guardrail_name="tool-call-scanner")],
|
||||
)
|
||||
|
||||
first_chunk = litellm.ModelResponseStream(
|
||||
id="chatcmpl-guardrail-tools",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call-shell",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "run_command",
|
||||
"arguments": '{"command":"blocked-',
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
created=0,
|
||||
model="gpt-4o-mini",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
second_chunk = litellm.ModelResponseStream(
|
||||
id="chatcmpl-guardrail-tools",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"function": {"arguments": 'command"}'},
|
||||
}
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
created=0,
|
||||
model="gpt-4o-mini",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
out = await proxy_logging.async_post_call_streaming_hook(
|
||||
data={},
|
||||
response=second_chunk,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
streaming_tool_calls_so_far=streaming_tool_calls_with_response((), first_chunk),
|
||||
)
|
||||
|
||||
assert out == "data: blocked\n\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_streaming_hook_preserves_equal_guardrail_projection(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
class _CopyingGuardrail(CustomGuardrail):
|
||||
def should_run_guardrail(self, data, event_type):
|
||||
return True
|
||||
|
||||
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
|
||||
return response.encode().decode()
|
||||
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"callbacks",
|
||||
[_CopyingGuardrail(guardrail_name="tool-call-scanner")],
|
||||
)
|
||||
response = litellm.ModelResponseStream(
|
||||
id="chatcmpl-guardrail-copy",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call-weather",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"city":"Paris"}',
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
created=0,
|
||||
model="gpt-4o-mini",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
out = await proxy_logging.async_post_call_streaming_hook(
|
||||
data={},
|
||||
response=response,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
)
|
||||
|
||||
assert out is response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_streaming_hook_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch):
|
||||
class _Per(CustomLogger):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue