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
This commit is contained in:
Krish Dholakia 2025-07-03 21:15:10 -07:00 • committed by GitHub
parent bffe5bec7d
commit f755b70528
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 77 additions and 25 deletions

View file

@ -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
)

View file

@ -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):

View file

@ -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

View file

@ -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