From b1aed699eaf17300b48233f066117390b8c41afb Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 15 Aug 2024 15:12:31 -0700 Subject: [PATCH] test sync sagemaker calls --- litellm/tests/test_sagemaker.py | 91 +++++++++++++++++++++++++++++---- 1 file changed, 81 insertions(+), 10 deletions(-) diff --git a/litellm/tests/test_sagemaker.py b/litellm/tests/test_sagemaker.py index 831ec5a2a8a..7293886cecd 100644 --- a/litellm/tests/test_sagemaker.py +++ b/litellm/tests/test_sagemaker.py @@ -44,19 +44,31 @@ def reset_callbacks(): @pytest.mark.asyncio() -async def test_completion_sagemaker(): +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_completion_sagemaker(sync_mode): try: litellm.set_verbose = True print("testing sagemaker") - response = await litellm.acompletion( - model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614", - messages=[ - {"role": "user", "content": "hi"}, - ], - temperature=0.2, - max_tokens=80, - input_cost_per_second=0.000420, - ) + if sync_mode is True: + response = litellm.completion( + model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614", + messages=[ + {"role": "user", "content": "hi"}, + ], + temperature=0.2, + max_tokens=80, + input_cost_per_second=0.000420, + ) + else: + response = await litellm.acompletion( + model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614", + messages=[ + {"role": "user", "content": "hi"}, + ], + temperature=0.2, + max_tokens=80, + input_cost_per_second=0.000420, + ) # Add any assertions here to check the response print(response) cost = completion_cost(completion_response=response) @@ -125,3 +137,62 @@ async def test_acompletion_sagemaker_non_stream(): kwargs["url"] == "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations" ) + + +@pytest.mark.asyncio +async def test_completion_sagemaker_non_stream(): + mock_response = MagicMock() + + def return_val(): + return { + "generated_text": "This is a mock response from SageMaker.", + "id": "cmpl-mockid", + "object": "text_completion", + "created": 1629800000, + "model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614", + "choices": [ + { + "text": "This is a mock response from SageMaker.", + "index": 0, + "logprobs": None, + "finish_reason": "length", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 8, "total_tokens": 9}, + } + + mock_response.json = return_val + + expected_payload = { + "inputs": "hi", + "parameters": {"temperature": 0.2, "max_new_tokens": 80}, + } + + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=mock_response, + ) as mock_post: + # Act: Call the litellm.acompletion function + response = litellm.completion( + model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614", + messages=[ + {"role": "user", "content": "hi"}, + ], + temperature=0.2, + max_tokens=80, + input_cost_per_second=0.000420, + ) + + # Print what was called on the mock + print("call args=", mock_post.call_args) + + # Assert + mock_post.assert_called_once() + _, kwargs = mock_post.call_args + args_to_sagemaker = kwargs["json"] + print("Arguments passed to sagemaker=", args_to_sagemaker) + assert args_to_sagemaker == expected_payload + assert ( + kwargs["url"] + == "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations" + )