mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-17 02:23:32 +00:00
Merge pull request #9326 from andjsmi/main
Modify completion handler for SageMaker to use payload from `prepared_request`
This commit is contained in:
@@ -213,7 +213,7 @@ class SagemakerLLM(BaseAWSLLM):
|
||||
sync_response = sync_handler.post(
|
||||
url=prepared_request.url,
|
||||
headers=prepared_request.headers, # type: ignore
|
||||
json=data,
|
||||
data=prepared_request.body,
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
@@ -308,7 +308,7 @@ class SagemakerLLM(BaseAWSLLM):
|
||||
sync_response = sync_handler.post(
|
||||
url=prepared_request.url,
|
||||
headers=prepared_request.headers, # type: ignore
|
||||
json=_data,
|
||||
data=prepared_request.body,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
@@ -356,7 +356,7 @@ class SagemakerLLM(BaseAWSLLM):
|
||||
self,
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: dict,
|
||||
data: str,
|
||||
logging_obj,
|
||||
client=None,
|
||||
):
|
||||
@@ -368,7 +368,7 @@ class SagemakerLLM(BaseAWSLLM):
|
||||
response = await client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
json=data,
|
||||
data=data,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
@@ -440,7 +440,7 @@ class SagemakerLLM(BaseAWSLLM):
|
||||
completion_stream = await self.make_async_call(
|
||||
api_base=prepared_request.url,
|
||||
headers=prepared_request.headers, # type: ignore
|
||||
data=data,
|
||||
data=prepared_request.body,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
streaming_response = CustomStreamWrapper(
|
||||
@@ -522,7 +522,7 @@ class SagemakerLLM(BaseAWSLLM):
|
||||
response = await async_handler.post(
|
||||
url=prepared_request.url,
|
||||
headers=prepared_request.headers, # type: ignore
|
||||
json=data,
|
||||
data=prepared_request.body,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
|
||||
@@ -265,7 +265,7 @@ async def test_acompletion_sagemaker_non_stream():
|
||||
# Assert
|
||||
mock_post.assert_called_once()
|
||||
_, kwargs = mock_post.call_args
|
||||
args_to_sagemaker = kwargs["json"]
|
||||
args_to_sagemaker = json.loads(kwargs["data"])
|
||||
print("Arguments passed to sagemaker=", args_to_sagemaker)
|
||||
assert args_to_sagemaker == expected_payload
|
||||
assert (
|
||||
@@ -325,7 +325,7 @@ async def test_completion_sagemaker_non_stream():
|
||||
# Assert
|
||||
mock_post.assert_called_once()
|
||||
_, kwargs = mock_post.call_args
|
||||
args_to_sagemaker = kwargs["json"]
|
||||
args_to_sagemaker = json.loads(kwargs["data"])
|
||||
print("Arguments passed to sagemaker=", args_to_sagemaker)
|
||||
assert args_to_sagemaker == expected_payload
|
||||
assert (
|
||||
@@ -386,7 +386,7 @@ async def test_completion_sagemaker_prompt_template_non_stream():
|
||||
# Assert
|
||||
mock_post.assert_called_once()
|
||||
_, kwargs = mock_post.call_args
|
||||
args_to_sagemaker = kwargs["json"]
|
||||
args_to_sagemaker = json.loads(kwargs["data"])
|
||||
print("Arguments passed to sagemaker=", args_to_sagemaker)
|
||||
assert args_to_sagemaker == expected_payload
|
||||
|
||||
@@ -445,7 +445,7 @@ async def test_completion_sagemaker_non_stream_with_aws_params():
|
||||
# Assert
|
||||
mock_post.assert_called_once()
|
||||
_, kwargs = mock_post.call_args
|
||||
args_to_sagemaker = kwargs["json"]
|
||||
args_to_sagemaker = json.loads(kwargs["data"])
|
||||
print("Arguments passed to sagemaker=", args_to_sagemaker)
|
||||
assert args_to_sagemaker == expected_payload
|
||||
assert (
|
||||
|
||||
Reference in New Issue
Block a user