Merge pull request #9326 from andjsmi/main

Modify completion handler for SageMaker to use payload from `prepared_request`
This commit is contained in:
Krish Dholakia 2025-03-17 22:16:43 -07:00 committed by GitHub
commit bcbb88d802
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 10 additions and 10 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

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