From 21f2ba6f1f4c115e9b35ddef6394b906ab5f7a2a Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 16 May 2024 23:20:51 -0700 Subject: [PATCH] fix(bedrock_httpx.py): logging fixes --- litellm/llms/bedrock_httpx.py | 31 ++++++++++++++++++++- litellm/tests/test_custom_callback_input.py | 2 +- 2 files changed, 31 insertions(+), 2 deletions(-) diff --git a/litellm/llms/bedrock_httpx.py b/litellm/llms/bedrock_httpx.py index 7085e58b328..df4cafe6e6a 100644 --- a/litellm/llms/bedrock_httpx.py +++ b/litellm/llms/bedrock_httpx.py @@ -735,7 +735,19 @@ class BedrockLLM(BaseLLM): inference_params[k] = v data = json.dumps({"prompt": prompt, **inference_params}) else: - raise Exception("UNSUPPORTED PROVIDER") + ## LOGGING + logging_obj.pre_call( + input=messages, + api_key="", + additional_args={ + "complete_input_dict": inference_params, + }, + ) + raise Exception( + "Bedrock HTTPX: Unsupported provider={}, model={}".format( + provider, model + ) + ) ## COMPLETION CALL @@ -822,6 +834,14 @@ class BedrockLLM(BaseLLM): status_code=response.status_code, message=response.text ) + ## LOGGING + logging_obj.post_call( + input=messages, + api_key="", + original_response=response.text, + additional_args={"complete_input_dict": data}, + ) + decoder = AWSEventStreamDecoder(model=model) completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=1024)) @@ -940,6 +960,15 @@ class BedrockLLM(BaseLLM): custom_llm_provider="bedrock", logging_obj=logging_obj, ) + + ## LOGGING + logging_obj.post_call( + input=messages, + api_key="", + original_response=streaming_response, + additional_args={"complete_input_dict": data}, + ) + return streaming_response def embedding(self, *args, **kwargs): diff --git a/litellm/tests/test_custom_callback_input.py b/litellm/tests/test_custom_callback_input.py index 2754ac65612..f4e16cdf35f 100644 --- a/litellm/tests/test_custom_callback_input.py +++ b/litellm/tests/test_custom_callback_input.py @@ -558,7 +558,7 @@ async def test_async_chat_bedrock_stream(): continue except: pass - time.sleep(1) + await asyncio.sleep(1) print(f"customHandler.errors: {customHandler.errors}") assert len(customHandler.errors) == 0 litellm.callbacks = []