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
2 changed files with 10 additions and 10 deletions
+6 -6
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,
)
+4 -4
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 (