fix(proxy): preserve tool calls through streaming hooks

This commit is contained in:
Ultronen 2026-09-05 17:54:01 +08:00 • committed by Ultronen
parent b3882d8e43
commit b5d1fb7298
3 changed files with 220 additions and 6 deletions

View file

@ -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

View file

@ -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)

View file

@ -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):