mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
bffe5bec7d
commit
f755b70528
4 changed files with 77 additions and 25 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue