From f755b70528d164227477382fb3831a4ea6d9283b Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Thu, 3 Jul 2025 21:15:10 -0700 Subject: [PATCH] fix(factory.py): support optional args for bedrock (#12287) * fix(factory.py): support optional args for bedrock Closes https://github.com/BerriAI/litellm/pull/12276 * test(main.py): Support async await on mock_delay Closes https://github.com/BerriAI/litellm/issues/12282 --- .../prompt_templates/factory.py | 2 +- litellm/main.py | 15 +++- litellm/utils.py | 12 ++- tests/test_litellm/test_main.py | 73 +++++++++++++------ 4 files changed, 77 insertions(+), 25 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index c99c0ae726e..2a1fc6cdc88 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -2631,7 +2631,7 @@ def _convert_to_bedrock_tool_call_invoke( id = tool["id"] name = tool["function"].get("name", "") arguments = tool["function"].get("arguments", "") - arguments_dict = json.loads(arguments) + arguments_dict = json.loads(arguments) if arguments else {} bedrock_tool = BedrockToolUseBlock( input=arguments_dict, name=name, toolUseId=id ) diff --git a/litellm/main.py b/litellm/main.py index 8fc2d8719d9..da6db8e8358 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -106,6 +106,7 @@ from litellm.utils import ( mock_completion_streaming_obj, pre_process_non_default_params, read_config_args, + should_run_mock_completion, supports_httpx_timeout, token_counter, validate_and_fix_openai_messages, @@ -451,6 +452,7 @@ async def acompletion( ######################################################### ######################################################### + # Adjusted to use explicit arguments instead of *args and **kwargs completion_kwargs = { "model": model, @@ -507,6 +509,15 @@ async def acompletion( ) return response + ### APPLY MOCK DELAY ### + + mock_delay = kwargs.get("mock_delay") + mock_response = kwargs.get("mock_response") + mock_tool_calls = kwargs.get("mock_tool_calls") + mock_timeout = kwargs.get("mock_timeout") + if mock_delay and should_run_mock_completion(mock_response=mock_response, mock_tool_calls=mock_tool_calls, mock_timeout=mock_timeout): + await asyncio.sleep(mock_delay) + try: # Use a partial function to pass your keyword arguments func = partial(completion, **completion_kwargs, **kwargs) @@ -673,6 +684,7 @@ async def _sleep_for_timeout_async(timeout: Union[float, str, httpx.Timeout]): await asyncio.sleep(timeout.connect) + def mock_completion( model: str, messages: List, @@ -710,6 +722,7 @@ def mock_completion( - If 'stream' is True, it returns a response that mimics the behavior of a streaming completion. """ try: + is_acompletion = kwargs.get("acompletion") or False if mock_response is None: mock_response = "This is a mock request" @@ -741,7 +754,7 @@ def mock_completion( status_code=529, ) time_delay = kwargs.get("mock_delay", None) - if time_delay is not None: + if time_delay is not None and not is_acompletion: time.sleep(time_delay) if isinstance(mock_response, dict): diff --git a/litellm/utils.py b/litellm/utils.py index 1aaea9fbafc..0f7c887ca91 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1405,7 +1405,7 @@ def client(original_function): # noqa: PLR0915 kwargs["max_tokens"] = modified_max_tokens except Exception as e: print_verbose(f"Error while checking max token limit: {str(e)}") - + # MODEL CALL result = await original_function(*args, **kwargs) end_time = datetime.datetime.now() @@ -7410,3 +7410,13 @@ def get_empty_usage() -> Usage: completion_tokens=0, total_tokens=0, ) + + +def should_run_mock_completion( + mock_response: Optional[Any], + mock_tool_calls: Optional[Any], + mock_timeout: Optional[Any], +) -> bool: + if mock_response or mock_tool_calls or mock_timeout: + return True + return False diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index c1e965dca3a..02f51dccf80 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -271,7 +271,9 @@ def test_bedrock_latency_optimized_inference(): def test_custom_provider_with_extra_headers(): from litellm.llms.custom_httpx.http_handler import HTTPHandler - with patch.object(litellm.llms.custom_httpx.http_handler.HTTPHandler, "post") as mock_post: + with patch.object( + litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" + ) as mock_post: response = litellm.completion( model="custom/custom", messages=[{"role": "user", "content": "Hello, how are you?"}], @@ -282,34 +284,42 @@ def test_custom_provider_with_extra_headers(): mock_post.assert_called_once() assert mock_post.call_args[1]["headers"]["X-Custom-Header"] == "custom-value" + def test_custom_provider_with_extra_body(): from litellm.llms.custom_httpx.http_handler import HTTPHandler - - with patch.object(litellm.llms.custom_httpx.http_handler.HTTPHandler, "post") as mock_post: + + with patch.object( + litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" + ) as mock_post: response = litellm.completion( model="custom/custom", messages=[{"role": "user", "content": "Hello, how are you?"}], - extra_body={"X-Custom-BodyValue": "custom-value", "X-Custom-BodyValue2": "custom-value2"}, + extra_body={ + "X-Custom-BodyValue": "custom-value", + "X-Custom-BodyValue2": "custom-value2", + }, api_base="https://example.com/api/v1", ) mock_post.assert_called_once() assert mock_post.call_args[1]["json"]["X-Custom-BodyValue"] == "custom-value" assert mock_post.call_args[1]["json"] == { - 'model': 'custom', - 'params': { - 'prompt': ['Hello, how are you?'], - 'max_tokens': None, - 'temperature': None, - 'top_p': None, - 'top_k': None + "model": "custom", + "params": { + "prompt": ["Hello, how are you?"], + "max_tokens": None, + "temperature": None, + "top_p": None, + "top_k": None, }, - 'X-Custom-BodyValue': 'custom-value', - 'X-Custom-BodyValue2': 'custom-value2' + "X-Custom-BodyValue": "custom-value", + "X-Custom-BodyValue2": "custom-value2", } # test that extra_body is not passed if not provided - with patch.object(litellm.llms.custom_httpx.http_handler.HTTPHandler, "post") as mock_post: + with patch.object( + litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" + ) as mock_post: response = litellm.completion( model="custom/custom", messages=[{"role": "user", "content": "Hello, how are you?"}], @@ -317,14 +327,14 @@ def test_custom_provider_with_extra_body(): ) mock_post.assert_called_once() assert mock_post.call_args[1]["json"] == { - 'model': 'custom', - 'params': { - 'prompt': ['Hello, how are you?'], - 'max_tokens': None, - 'temperature': None, - 'top_p': None, - 'top_k': None - } + "model": "custom", + "params": { + "prompt": ["Hello, how are you?"], + "max_tokens": None, + "temperature": None, + "top_p": None, + "top_k": None, + }, } @@ -518,3 +528,22 @@ def test_responses_api_bridge_check_handles_exception(): assert model == "custom-model" assert model_info["mode"] == "responses" + + +@pytest.mark.asyncio +async def test_async_mock_delay(): + """Use asyncio await for mock delay on acompletion""" + import time + + from litellm import acompletion + + start_time = time.time() + result = await acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Hey, how's it going?"}], + mock_delay=0.01, + mock_response="Hello world", + ) + end_time = time.time() + delay = end_time - start_time + assert delay >= 0.01