fix(bedrock_httpx.py): Fix https://github.com/BerriAI/litellm/issues/4415
This commit is contained in:
parent
1821b32491
commit
d1cb4a195c
@ -1,3 +1,8 @@
|
||||
####################################
|
||||
######### DEPRECATED FILE ##########
|
||||
####################################
|
||||
# logic moved to `bedrock_httpx.py` #
|
||||
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
|
||||
@ -261,20 +261,24 @@ class BedrockLLM(BaseLLM):
|
||||
# handle anthropic prompts and amazon titan prompts
|
||||
prompt = ""
|
||||
chat_history: Optional[list] = None
|
||||
## CUSTOM PROMPT
|
||||
if model in custom_prompt_dict:
|
||||
# check if the model has a registered custom prompt
|
||||
model_prompt_details = custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details.get(
|
||||
"initial_prompt_value", ""
|
||||
),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
return prompt, None
|
||||
## ELSE
|
||||
if provider == "anthropic" or provider == "amazon":
|
||||
if model in custom_prompt_dict:
|
||||
# check if the model has a registered custom prompt
|
||||
model_prompt_details = custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
messages=messages,
|
||||
)
|
||||
else:
|
||||
prompt = prompt_factory(
|
||||
model=model, messages=messages, custom_llm_provider="bedrock"
|
||||
)
|
||||
prompt = prompt_factory(
|
||||
model=model, messages=messages, custom_llm_provider="bedrock"
|
||||
)
|
||||
elif provider == "mistral":
|
||||
prompt = prompt_factory(
|
||||
model=model, messages=messages, custom_llm_provider="bedrock"
|
||||
|
||||
@ -1,20 +1,31 @@
|
||||
# @pytest.mark.skip(reason="AWS Suspended Account")
|
||||
import sys, os
|
||||
import os
|
||||
import sys
|
||||
import traceback
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
import os, io
|
||||
import io
|
||||
import os
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import embedding, completion, completion_cost, Timeout, ModelResponse
|
||||
from litellm import RateLimitError
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
|
||||
from unittest.mock import patch, AsyncMock, Mock
|
||||
from litellm import (
|
||||
ModelResponse,
|
||||
RateLimitError,
|
||||
Timeout,
|
||||
completion,
|
||||
completion_cost,
|
||||
embedding,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
# litellm.num_retries = 3
|
||||
litellm.cache = None
|
||||
@ -481,7 +492,10 @@ def test_completion_claude_3_base64():
|
||||
def test_provisioned_throughput():
|
||||
try:
|
||||
litellm.set_verbose = True
|
||||
import botocore, json, io
|
||||
import io
|
||||
import json
|
||||
|
||||
import botocore
|
||||
import botocore.session
|
||||
from botocore.stub import Stubber
|
||||
|
||||
@ -537,7 +551,6 @@ def test_completion_bedrock_mistral_completion_auth():
|
||||
# aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"]
|
||||
# aws_secret_access_key = os.environ["AWS_SECRET_ACCESS_KEY"]
|
||||
# aws_region_name = os.environ["AWS_REGION_NAME"]
|
||||
|
||||
# os.environ.pop("AWS_ACCESS_KEY_ID", None)
|
||||
# os.environ.pop("AWS_SECRET_ACCESS_KEY", None)
|
||||
# os.environ.pop("AWS_REGION_NAME", None)
|
||||
@ -624,3 +637,48 @@ async def test_bedrock_extra_headers():
|
||||
assert "test" in mock_client_post.call_args.kwargs["headers"]
|
||||
assert mock_client_post.call_args.kwargs["headers"]["test"] == "hello world"
|
||||
mock_client_post.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_custom_prompt_template():
|
||||
"""
|
||||
Check if custom prompt template used for bedrock models
|
||||
|
||||
Reference: https://github.com/BerriAI/litellm/issues/4415
|
||||
"""
|
||||
client = AsyncHTTPHandler()
|
||||
|
||||
with patch.object(client, "post", new=AsyncMock()) as mock_client_post:
|
||||
import json
|
||||
|
||||
try:
|
||||
response = await litellm.acompletion(
|
||||
model="bedrock/mistral.OpenOrca",
|
||||
messages=[{"role": "user", "content": "What's AWS?"}],
|
||||
client=client,
|
||||
roles={
|
||||
"system": {
|
||||
"pre_message": "<|im_start|>system\n",
|
||||
"post_message": "<|im_end|>",
|
||||
},
|
||||
"assistant": {
|
||||
"pre_message": "<|im_start|>assistant\n",
|
||||
"post_message": "<|im_end|>",
|
||||
},
|
||||
"user": {
|
||||
"pre_message": "<|im_start|>user\n",
|
||||
"post_message": "<|im_end|>",
|
||||
},
|
||||
},
|
||||
bos_token="<s>",
|
||||
eos_token="<|im_end|>",
|
||||
)
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
print(f"mock_client_post.call_args: {mock_client_post.call_args}")
|
||||
assert "prompt" in mock_client_post.call_args.kwargs["data"]
|
||||
|
||||
prompt = json.loads(mock_client_post.call_args.kwargs["data"])["prompt"]
|
||||
assert prompt == "<|im_start|>user\nWhat's AWS?<|im_end|>"
|
||||
mock_client_post.assert_called_once()
|
||||
|
||||
Loading…
Reference in New Issue
Block a user