From 9bfbdc18fbaf9572604d9074cb2557d604be5082 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 9 Nov 2023 09:17:43 -0800 Subject: [PATCH] feat(utils.py): enable returning complete response when stream=true --- litellm/main.py | 2 +- litellm/tests/test_stream_chunk_builder.py | 4 +++- litellm/tests/test_streaming.py | 11 +++++++++++ litellm/utils.py | 12 +++++++----- 4 files changed, 22 insertions(+), 7 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 2d6d200b8f5..21eb580e998 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -270,7 +270,7 @@ def completion( eos_token = kwargs.get("eos_token", None) ######## end of unpacking kwargs ########### openai_params = ["functions", "function_call", "temperature", "temperature", "top_p", "n", "stream", "stop", "max_tokens", "presence_penalty", "frequency_penalty", "logit_bias", "user", "request_timeout", "api_base", "api_version", "api_key"] - litellm_params = ["metadata", "acompletion", "caching", "return_async", "mock_response", "api_key", "api_version", "api_base", "force_timeout", "logger_fn", "verbose", "custom_llm_provider", "litellm_logging_obj", "litellm_call_id", "use_client", "id", "fallbacks", "azure", "headers", "model_list", "num_retries", "context_window_fallback_dict", "roles", "final_prompt_value", "bos_token", "eos_token", "request_timeout"] + litellm_params = ["metadata", "acompletion", "caching", "return_async", "mock_response", "api_key", "api_version", "api_base", "force_timeout", "logger_fn", "verbose", "custom_llm_provider", "litellm_logging_obj", "litellm_call_id", "use_client", "id", "fallbacks", "azure", "headers", "model_list", "num_retries", "context_window_fallback_dict", "roles", "final_prompt_value", "bos_token", "eos_token", "request_timeout", "complete_response"] default_params = openai_params + litellm_params non_default_params = {k: v for k,v in kwargs.items() if k not in default_params} # model-specific params - pass them straight to the model/provider if mock_response: diff --git a/litellm/tests/test_stream_chunk_builder.py b/litellm/tests/test_stream_chunk_builder.py index 3a4d25b58c3..7b8db048104 100644 --- a/litellm/tests/test_stream_chunk_builder.py +++ b/litellm/tests/test_stream_chunk_builder.py @@ -24,6 +24,7 @@ function_schema = { } def test_stream_chunk_builder(): + litellm.set_verbose = False litellm.api_key = os.environ["OPENAI_API_KEY"] response = completion( model="gpt-3.5-turbo", @@ -35,10 +36,11 @@ def test_stream_chunk_builder(): chunks = [] for chunk in response: - print(chunk) + # print(chunk) chunks.append(chunk) try: + print(f"chunks: {chunks}") rebuilt_response = stream_chunk_builder(chunks) # exract the response from the rebuilt response diff --git a/litellm/tests/test_streaming.py b/litellm/tests/test_streaming.py index 179f7f151c6..bb899b25957 100644 --- a/litellm/tests/test_streaming.py +++ b/litellm/tests/test_streaming.py @@ -902,6 +902,17 @@ def test_openai_chat_completion_call(): # test_openai_chat_completion_call() +def test_openai_chat_completion_complete_response_call(): + try: + complete_response = completion( + model="gpt-3.5-turbo", messages=messages, stream=True, complete_response=True + ) + print(f"complete response: {complete_response}") + except: + print(f"error occurred: {traceback.format_exc()}") + pass + +test_openai_chat_completion_complete_response_call() def test_openai_text_completion_call(): try: diff --git a/litellm/utils.py b/litellm/utils.py index 6b83138fe50..a75822ff079 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -949,16 +949,18 @@ def client(original_function): end_time = datetime.datetime.now() if "stream" in kwargs and kwargs["stream"] == True: # TODO: Add to cache for streaming - return result + if "complete_response" in kwargs and kwargs["complete_response"] == True: + chunks = [] + for idx, chunk in enumerate(result): + chunks.append(chunk) + return litellm.stream_chunk_builder(chunks) + else: + return result # [OPTIONAL] ADD TO CACHE if litellm.caching or litellm.caching_with_models or litellm.cache != None: # user init a cache object litellm.cache.add_cache(result, *args, **kwargs) - - # [OPTIONAL] Return LiteLLM call_id - if litellm.use_client == True: - result['litellm_call_id'] = litellm_call_id # LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated logging_obj.success_handler(result, start_time, end_time)