From bd37a9cb5e3e34aa4550a41c43ab9c28db825f6f Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Tue, 23 Jan 2024 11:12:16 -0800 Subject: [PATCH 1/4] (fix) proxy - streaming sagemaker --- litellm/proxy/proxy_server.py | 26 ++++++++++++++++++-------- litellm/proxy/tests/test_openai_js.js | 2 +- 2 files changed, 19 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 78e756a2a6a..f4eb04874b6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1658,11 +1658,16 @@ async def completion( "stream" in data and data["stream"] == True ): # use generate_responses to stream responses custom_headers = {"x-litellm-model-id": model_id} - return StreamingResponse( - async_data_generator( - user_api_key_dict=user_api_key_dict, + stream_content = async_data_generator( + user_api_key_dict=user_api_key_dict, + response=response, + ) + if response.custom_llm_provider == "sagemaker": + stream_content = data_generator( response=response, - ), + ) + return StreamingResponse( + stream_content, media_type="text/event-stream", headers=custom_headers, ) @@ -1820,11 +1825,16 @@ async def chat_completion( "stream" in data and data["stream"] == True ): # use generate_responses to stream responses custom_headers = {"x-litellm-model-id": model_id} - return StreamingResponse( - async_data_generator( - user_api_key_dict=user_api_key_dict, + stream_content = async_data_generator( + user_api_key_dict=user_api_key_dict, + response=response, + ) + if response.custom_llm_provider == "sagemaker": + stream_content = data_generator( response=response, - ), + ) + return StreamingResponse( + stream_content, media_type="text/event-stream", headers=custom_headers, ) diff --git a/litellm/proxy/tests/test_openai_js.js b/litellm/proxy/tests/test_openai_js.js index 7e74eeca3f8..c0f25cf0585 100644 --- a/litellm/proxy/tests/test_openai_js.js +++ b/litellm/proxy/tests/test_openai_js.js @@ -4,7 +4,7 @@ const openai = require('openai'); process.env.DEBUG=false; async function runOpenAI() { const client = new openai.OpenAI({ - apiKey: 'sk-yPX56TDqBpr23W7ruFG3Yg', + apiKey: 'sk-JkKeNi6WpWDngBsghJ6B9g', baseURL: 'http://0.0.0.0:8000' }); From a61dbc1613c4696a9d6a0d675371a6a7d21a5974 Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Tue, 23 Jan 2024 12:08:58 -0800 Subject: [PATCH 2/4] (fix) select_data_generator - sagemaker --- litellm/proxy/proxy_server.py | 37 ++++++++++++++++++++--------------- 1 file changed, 21 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f4eb04874b6..af5d6d5ac9f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1422,6 +1422,19 @@ async def async_data_generator(response, user_api_key_dict): yield f"data: {str(e)}\n\n" +def select_data_generator(response, user_api_key_dict): + # since boto3 - sagemaker does not support async calls + if response.custom_llm_provider == "sagemaker": + return data_generator( + response=response, + ) + else: + # default to async_data_generator + return async_data_generator( + response=response, user_api_key_dict=user_api_key_dict + ) + + def get_litellm_model_info(model: dict = {}): model_info = model.get("model_info", {}) model_to_lookup = model.get("litellm_params", {}).get("model", None) @@ -1658,16 +1671,12 @@ async def completion( "stream" in data and data["stream"] == True ): # use generate_responses to stream responses custom_headers = {"x-litellm-model-id": model_id} - stream_content = async_data_generator( - user_api_key_dict=user_api_key_dict, - response=response, + selected_data_generator = select_data_generator( + response=response, user_api_key_dict=user_api_key_dict ) - if response.custom_llm_provider == "sagemaker": - stream_content = data_generator( - response=response, - ) + return StreamingResponse( - stream_content, + selected_data_generator, media_type="text/event-stream", headers=custom_headers, ) @@ -1825,16 +1834,12 @@ async def chat_completion( "stream" in data and data["stream"] == True ): # use generate_responses to stream responses custom_headers = {"x-litellm-model-id": model_id} - stream_content = async_data_generator( - user_api_key_dict=user_api_key_dict, - response=response, + selected_data_generator = select_data_generator( + response=response, user_api_key_dict=user_api_key_dict ) - if response.custom_llm_provider == "sagemaker": - stream_content = data_generator( - response=response, - ) + return StreamingResponse( - stream_content, + selected_data_generator, media_type="text/event-stream", headers=custom_headers, ) From 44e213e842d0ad7121aa56f39b156de6028686ca Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Tue, 23 Jan 2024 12:13:34 -0800 Subject: [PATCH 3/4] (fix) select_data_generator --- litellm/proxy/proxy_server.py | 23 ++++++++++++++++------- 1 file changed, 16 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index af5d6d5ac9f..a1790f49cea 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1423,13 +1423,22 @@ async def async_data_generator(response, user_api_key_dict): def select_data_generator(response, user_api_key_dict): - # since boto3 - sagemaker does not support async calls - if response.custom_llm_provider == "sagemaker": - return data_generator( - response=response, - ) - else: - # default to async_data_generator + try: + # since boto3 - sagemaker does not support async calls, we should use a sync data_generator + if ( + hasattr(response, "custom_llm_provider") + and response.custom_llm_provider == "sagemaker" + ): + return data_generator( + response=response, + ) + else: + # default to async_data_generator + return async_data_generator( + response=response, user_api_key_dict=user_api_key_dict + ) + except: + # worst case - use async_data_generator return async_data_generator( response=response, user_api_key_dict=user_api_key_dict ) From e8cd27f2b75277cd8ae2f2ea02014ac6e4e1ecfd Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Tue, 23 Jan 2024 12:31:16 -0800 Subject: [PATCH 4/4] (fix) sagemaker streaming support --- litellm/utils.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/litellm/utils.py b/litellm/utils.py index 00b76bfb5e4..85d160334ea 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7732,10 +7732,8 @@ class CustomStreamWrapper: ] self.sent_last_chunk = True elif self.custom_llm_provider == "sagemaker": - print_verbose(f"ENTERS SAGEMAKER STREAMING") - new_chunk = next(self.completion_stream) - - completion_obj["content"] = new_chunk + print_verbose(f"ENTERS SAGEMAKER STREAMING for chunk {chunk}") + completion_obj["content"] = chunk elif self.custom_llm_provider == "petals": if len(self.completion_stream) == 0: if self.sent_last_chunk: @@ -7854,7 +7852,7 @@ class CustomStreamWrapper: completion_obj["role"] = "assistant" self.sent_first_chunk = True model_response.choices[0].delta = Delta(**completion_obj) - print_verbose(f"model_response: {model_response}") + print_verbose(f"returning model_response: {model_response}") return model_response else: return