mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(utils): run post-call deployment hook on converted chat streams
Backport of #41495 to rc/1.102.0.
Cherry-picked from merge commit 0add8c0083 (main), originally by app/devin-ai-integration.
The only conflict was the test import header: the rc line never gained the logging_executor import that main carries next to this change, so only the MockResponseIterator import comes along.
This commit is contained in:
parent
8f8b47d6f0
commit
707a81aeae
2 changed files with 134 additions and 1 deletions
|
|
@ -851,6 +851,32 @@ def _is_converted_stream_result(result: object) -> bool:
|
|||
return isinstance(result, (CustomStreamWrapper, BaseResponsesAPIStreamingIterator))
|
||||
|
||||
|
||||
async def _run_success_deployment_hook_on_converted_chat_stream(
|
||||
result: object, request_data: dict[str, object], call_type: str
|
||||
) -> None:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
|
||||
if not isinstance(result, CustomStreamWrapper):
|
||||
return
|
||||
completion_stream: Final = result.completion_stream
|
||||
if not isinstance(completion_stream, MockResponseIterator):
|
||||
return
|
||||
call_type_enum: Final = _CALL_TYPE_ENUM_MAP.get(call_type)
|
||||
if call_type_enum is None:
|
||||
return
|
||||
hooked: Final = await async_post_call_success_deployment_hook(
|
||||
request_data=request_data,
|
||||
response=completion_stream.model_response,
|
||||
call_type=call_type_enum,
|
||||
)
|
||||
if not isinstance(hooked, ModelResponse) or hooked is completion_stream.model_response:
|
||||
return
|
||||
result.completion_stream = MockResponseIterator( # rebind-ok: a new wrapper would drop headers and fire __del__
|
||||
model_response=hooked, json_mode=completion_stream.json_mode
|
||||
)
|
||||
|
||||
|
||||
# Runs once per call to check if the user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
|
||||
def function_setup(
|
||||
original_function: str,
|
||||
|
|
@ -1949,9 +1975,14 @@ def client(original_function):
|
|||
raise
|
||||
end_time = datetime.datetime.now()
|
||||
|
||||
if _is_streaming_request(kwargs=kwargs, call_type=call_type) or _is_converted_stream_result(result):
|
||||
streaming_requested: Final = _is_streaming_request(kwargs=kwargs, call_type=call_type)
|
||||
if streaming_requested or _is_converted_stream_result(result):
|
||||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["stream"] = True
|
||||
if not streaming_requested:
|
||||
await _run_success_deployment_hook_on_converted_chat_stream(
|
||||
result=result, request_data=kwargs, call_type=call_type
|
||||
)
|
||||
if "complete_response" in kwargs and kwargs["complete_response"] is True:
|
||||
chunks: Final = []
|
||||
for idx, chunk in enumerate(result):
|
||||
|
|
|
|||
|
|
@ -32,12 +32,15 @@ from litellm._logging import (
|
|||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.proxy.utils import is_valid_api_key
|
||||
from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
Choices,
|
||||
Delta,
|
||||
LlmProviders,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
PromptTokensDetailsWrapper,
|
||||
StreamingChoices,
|
||||
|
|
@ -52,6 +55,7 @@ from litellm.utils import (
|
|||
_check_provider_match,
|
||||
_get_potential_model_names,
|
||||
_is_streaming_request,
|
||||
_run_success_deployment_hook_on_converted_chat_stream,
|
||||
_snapshot_exception_for_hook,
|
||||
async_post_call_failure_deployment_hook,
|
||||
async_post_call_success_deployment_hook,
|
||||
|
|
@ -5254,6 +5258,104 @@ async def test_wrapper_async_logs_converted_chat_stream_with_standard_logging_ob
|
|||
assert success_kwargs["stream"] is True
|
||||
|
||||
|
||||
class _RewritingSuccessDeploymentHook(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.seen_responses: tuple[object, ...] = ()
|
||||
|
||||
async def async_post_call_success_deployment_hook(
|
||||
self, request_data: dict[str, object], response: object, call_type: CallTypes | None
|
||||
) -> ModelResponse | None:
|
||||
self.seen_responses = (*self.seen_responses, response)
|
||||
if not isinstance(response, ModelResponse):
|
||||
return None
|
||||
choice: Final = response.choices[0]
|
||||
if not isinstance(choice, Choices):
|
||||
return None
|
||||
rewritten_message: Final = choice.message.model_copy(update={"content": "rewritten by deployment hook"})
|
||||
return response.model_copy(update={"choices": [choice.model_copy(update={"message": rewritten_message})]})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_async_runs_success_deployment_hook_on_converted_chat_stream(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_install_converted_stream_callbacks(monkeypatch)
|
||||
hook: Final = _RewritingSuccessDeploymentHook()
|
||||
monkeypatch.setattr(litellm, "callbacks", [_ConvertStreamDeploymentHook(), hook])
|
||||
|
||||
response: Final = await litellm.acompletion(
|
||||
model="gpt-5.6",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
mock_response="converted stream body",
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, CustomStreamWrapper)
|
||||
chunks: Final = [chunk async for chunk in response]
|
||||
|
||||
assert len(hook.seen_responses) == 1
|
||||
seen: Final = hook.seen_responses[0]
|
||||
assert isinstance(seen, ModelResponse)
|
||||
assert seen.choices[0].message.content == "converted stream body"
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "rewritten by deployment hook"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("completion_stream", "call_type"),
|
||||
[
|
||||
(iter([ModelResponse(model="gpt-5.6")]), "acompletion"),
|
||||
(MockResponseIterator(model_response=ModelResponse(model="gpt-5.6")), "not_a_call_type"),
|
||||
],
|
||||
ids=["real_provider_stream", "unmapped_call_type"],
|
||||
)
|
||||
async def test_converted_chat_stream_hook_skips_unhandled_wrappers(
|
||||
monkeypatch: pytest.MonkeyPatch, completion_stream: object, call_type: str
|
||||
) -> None:
|
||||
hook: Final = _RewritingSuccessDeploymentHook()
|
||||
monkeypatch.setattr(litellm, "callbacks", [hook])
|
||||
wrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream, model="gpt-5.6", logging_obj=MagicMock(), custom_llm_provider="openai"
|
||||
)
|
||||
|
||||
await _run_success_deployment_hook_on_converted_chat_stream(
|
||||
result=wrapper, request_data={"model": "gpt-5.6"}, call_type=call_type
|
||||
)
|
||||
|
||||
assert hook.seen_responses == ()
|
||||
assert wrapper.completion_stream is completion_stream
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
async def test_wrapper_async_leaves_success_deployment_hook_off_requested_fake_stream(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
hook: Final = _RewritingSuccessDeploymentHook()
|
||||
monkeypatch.setattr(litellm, "callbacks", [hook])
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
respx.post("http://fake-stream.invalid/api/v1/run/flow-1").respond(
|
||||
json={"outputs": [{"outputs": [{"results": {"message": {"text": "plain stream body"}}}]}]}
|
||||
)
|
||||
|
||||
response: Final = await litellm.acompletion(
|
||||
model="langflow/flow-1",
|
||||
api_base="http://fake-stream.invalid",
|
||||
api_key="fake-key",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, CustomStreamWrapper)
|
||||
assert isinstance(response.completion_stream, MockResponseIterator)
|
||||
chunks: Final = [chunk async for chunk in response]
|
||||
|
||||
assert hook.seen_responses == ()
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "plain stream body"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
async def test_wrapper_async_logs_converted_responses_stream_with_standard_logging_object(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue