Merge pull request #9335 from BerriAI/litellm_dev_03_17_2025_p3

Contributor PR: Fix sagemaker too little data for content error
This commit is contained in:
Krish Dholakia 2025-03-18 23:24:07 -07:00 • committed by GitHub
commit 6347b694ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 15 additions and 14 deletions

View file

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

View file

@ -5,7 +5,7 @@
import os
import sys
import traceback
import json
import pytest
sys.path.insert(
@ -465,7 +465,8 @@ def test_sagemaker_default_region():
)
mock_post.assert_called_once()
_, kwargs = mock_post.call_args
args_to_sagemaker = kwargs["json"]
print(f"kwargs: {kwargs}")
args_to_sagemaker = json.loads(kwargs["data"])
print("Arguments passed to sagemaker=", args_to_sagemaker)
print("url=", kwargs["url"])
@ -517,7 +518,7 @@ def test_sagemaker_environment_region():
)
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)
print("url=", kwargs["url"])
@ -574,7 +575,7 @@ def test_sagemaker_config_region():
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)
print("url=", kwargs["url"])

View file

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