mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): preserve tool calls through streaming hooks
This commit is contained in:
parent
b3882d8e43
commit
b5d1fb7298
3 changed files with 220 additions and 6 deletions
|
|
@ -892,6 +892,27 @@ 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."""
|
||||
|
||||
|
||||
def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, response: object) -> str:
|
||||
complete_response = 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 _is_unchanged_structured_streaming_hook_response(
|
||||
*, callback_response: object, complete_response: str, response_str: str, response: object
|
||||
) -> bool:
|
||||
if response_str != "" or not isinstance(response, (ModelResponse, ModelResponseStream)):
|
||||
return False
|
||||
if isinstance(complete_response, _StreamingHookResponseText):
|
||||
return callback_response is complete_response
|
||||
return callback_response == complete_response
|
||||
|
||||
|
||||
def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""Failure-path callbacks run after ``litellm_logging_obj`` is popped from
|
||||
request_data (it is not serialisable), so the caller merges these fields
|
||||
|
|
@ -3513,18 +3534,29 @@ 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,
|
||||
)
|
||||
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:
|
||||
# A text result cannot represent a structured empty-text
|
||||
# chunk such as a tool-call delta. Preserve the chunk
|
||||
# only when the callback returned its input unchanged.
|
||||
if _is_unchanged_structured_streaming_hook_response(
|
||||
callback_response=callback_response,
|
||||
complete_response=complete_response,
|
||||
response_str=response_str,
|
||||
response=response,
|
||||
):
|
||||
continue
|
||||
response = callback_response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -24,8 +24,9 @@ from fastapi import Response
|
|||
from fastapi.responses import StreamingResponse
|
||||
|
||||
import litellm
|
||||
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
|
||||
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,
|
||||
|
|
@ -682,6 +683,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)
|
||||
|
|
|
|||
|
|
@ -257,6 +257,113 @@ 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_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch):
|
||||
class _Per(CustomLogger):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue