From 7d0086d742ed6666ea7b8251a3dcef49dbafe069 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 17 Apr 2024 17:43:41 -0700 Subject: [PATCH 1/3] fix(utils.py): ensure streaming output parsing only applied for hf / sagemaker models selectively applies the checking --- litellm/tests/test_streaming.py | 14 ++++++++++++++ litellm/utils.py | 9 +++++++++ 2 files changed, 23 insertions(+) diff --git a/litellm/tests/test_streaming.py b/litellm/tests/test_streaming.py index a18da18192f..aa2a91b9f32 100644 --- a/litellm/tests/test_streaming.py +++ b/litellm/tests/test_streaming.py @@ -220,6 +220,20 @@ tools_schema = [ # test_completion_cohere_stream() +def test_completion_azure_stream_special_char(): + messages = [ + {"role": "user", "content": "Respond with the '<' sign and nothing else."} + ] + response = completion(model="azure/chatgpt-v-2", messages=messages, stream=True) + response_str = "" + for part in response: + response_str += part.choices[0].delta.content or "" + + print(f"response_str: {response_str}") + assert len(response_str) > 0 + raise Exception("it worked") + + def test_completion_cohere_stream_bad_key(): try: litellm.cache = None diff --git a/litellm/utils.py b/litellm/utils.py index 2ae1467d076..bea24c02fef 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8856,7 +8856,16 @@ class CustomStreamWrapper: raise e def check_special_tokens(self, chunk: str, finish_reason: Optional[str]): + """ + Output parse / special tokens for sagemaker + hf streaming. + """ hold = False + if ( + self.custom_llm_provider != "huggingface" + and self.custom_llm_provider != "sagemaker" + ): + return hold, chunk + if finish_reason: for token in self.special_tokens: if token in chunk: From 15ae7a8314fec6bb1e7d77933efd442de39eac33 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 17 Apr 2024 18:03:40 -0700 Subject: [PATCH 2/3] fix(utils.py): fix streaming special character flushing logic --- litellm/tests/test_streaming.py | 3 +-- litellm/utils.py | 13 +++++++------ 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/litellm/tests/test_streaming.py b/litellm/tests/test_streaming.py index aa2a91b9f32..32976978231 100644 --- a/litellm/tests/test_streaming.py +++ b/litellm/tests/test_streaming.py @@ -221,6 +221,7 @@ tools_schema = [ def test_completion_azure_stream_special_char(): + litellm.set_verbose = True messages = [ {"role": "user", "content": "Respond with the '<' sign and nothing else."} ] @@ -229,9 +230,7 @@ def test_completion_azure_stream_special_char(): for part in response: response_str += part.choices[0].delta.content or "" - print(f"response_str: {response_str}") assert len(response_str) > 0 - raise Exception("it worked") def test_completion_cohere_stream_bad_key(): diff --git a/litellm/utils.py b/litellm/utils.py index bea24c02fef..3e31874bd04 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8860,11 +8860,11 @@ class CustomStreamWrapper: Output parse / special tokens for sagemaker + hf streaming. """ hold = False - if ( - self.custom_llm_provider != "huggingface" - and self.custom_llm_provider != "sagemaker" - ): - return hold, chunk + # if ( + # self.custom_llm_provider != "huggingface" + # and self.custom_llm_provider != "sagemaker" + # ): + # return hold, chunk if finish_reason: for token in self.special_tokens: @@ -8881,6 +8881,7 @@ class CustomStreamWrapper: for token in self.special_tokens: if len(curr_chunk) < len(token) and curr_chunk in token: hold = True + self.holding_chunk = curr_chunk elif len(curr_chunk) >= len(token): if token in curr_chunk: self.holding_chunk = curr_chunk.replace(token, "") @@ -9962,6 +9963,7 @@ class CustomStreamWrapper: f"model_response.choices[0].delta: {model_response.choices[0].delta}; completion_obj: {completion_obj}" ) print_verbose(f"self.sent_first_chunk: {self.sent_first_chunk}") + ## RETURN ARG if ( "content" in completion_obj @@ -10034,7 +10036,6 @@ class CustomStreamWrapper: elif self.received_finish_reason is not None: if self.sent_last_chunk == True: raise StopIteration - # flush any remaining holding chunk if len(self.holding_chunk) > 0: if model_response.choices[0].delta.content is None: From 0d2b400e91042c029a231b7d82c94622d62c99e3 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 17 Apr 2024 18:32:34 -0700 Subject: [PATCH 3/3] test(test_function_calling.py): handle for when model returns a text response --- litellm/tests/test_function_calling.py | 79 +++++++++++++------------- 1 file changed, 41 insertions(+), 38 deletions(-) diff --git a/litellm/tests/test_function_calling.py b/litellm/tests/test_function_calling.py index 2a815edc8e0..f76a082f60e 100644 --- a/litellm/tests/test_function_calling.py +++ b/litellm/tests/test_function_calling.py @@ -266,47 +266,50 @@ def test_groq_parallel_function_call(): ) print("Response\n", response) response_message = response.choices[0].message - tool_calls = response_message.tool_calls + if hasattr(response_message, "tool_calls"): + tool_calls = response_message.tool_calls - assert isinstance(response.choices[0].message.tool_calls[0].function.name, str) - assert isinstance( - response.choices[0].message.tool_calls[0].function.arguments, str - ) + assert isinstance( + response.choices[0].message.tool_calls[0].function.name, str + ) + assert isinstance( + response.choices[0].message.tool_calls[0].function.arguments, str + ) - print("length of tool calls", len(tool_calls)) + print("length of tool calls", len(tool_calls)) - # Step 2: check if the model wanted to call a function - if tool_calls: - # Step 3: call the function - # Note: the JSON response may not always be valid; be sure to handle errors - available_functions = { - "get_current_weather": get_current_weather, - } # only one function in this example, but you can have multiple - messages.append( - response_message - ) # extend conversation with assistant's reply - print("Response message\n", response_message) - # Step 4: send the info for each function call and function response to the model - for tool_call in tool_calls: - function_name = tool_call.function.name - function_to_call = available_functions[function_name] - function_args = json.loads(tool_call.function.arguments) - function_response = function_to_call( - location=function_args.get("location"), - unit=function_args.get("unit"), - ) + # Step 2: check if the model wanted to call a function + if tool_calls: + # Step 3: call the function + # Note: the JSON response may not always be valid; be sure to handle errors + available_functions = { + "get_current_weather": get_current_weather, + } # only one function in this example, but you can have multiple messages.append( - { - "tool_call_id": tool_call.id, - "role": "tool", - "name": function_name, - "content": function_response, - } - ) # extend conversation with function response - print(f"messages: {messages}") - second_response = litellm.completion( - model="groq/llama2-70b-4096", messages=messages - ) # get a new response from the model where it can see the function response - print("second response\n", second_response) + response_message + ) # extend conversation with assistant's reply + print("Response message\n", response_message) + # Step 4: send the info for each function call and function response to the model + for tool_call in tool_calls: + function_name = tool_call.function.name + function_to_call = available_functions[function_name] + function_args = json.loads(tool_call.function.arguments) + function_response = function_to_call( + location=function_args.get("location"), + unit=function_args.get("unit"), + ) + messages.append( + { + "tool_call_id": tool_call.id, + "role": "tool", + "name": function_name, + "content": function_response, + } + ) # extend conversation with function response + print(f"messages: {messages}") + second_response = litellm.completion( + model="groq/llama2-70b-4096", messages=messages + ) # get a new response from the model where it can see the function response + print("second response\n", second_response) except Exception as e: pytest.fail(f"Error occurred: {e}")