From a907a85c342505acff4ec34ddf1b0a02bde139b6 Mon Sep 17 00:00:00 2001 From: Lucas Jia Date: Fri, 25 Sep 2026 13:23:04 -0700 Subject: [PATCH] fix: Keep AsyncPredictor serializers in sync with wrapped predictor AsyncPredictor copied the wrapped Predictor's serializer and deserializer onto itself at construction time. Requests were then serialized with the outer copy, while the Accept header and response decoding used the inner predictor's deserializer. Overriding AsyncPredictor.deserializer after creation therefore had no effect. Make serializer and deserializer properties that read from and write to the wrapped Predictor, so a single source of truth drives the upload, the Accept header, and result decoding. Fixes #3100 --- src/sagemaker/predictor_async.py | 30 +++++++++++++- tests/unit/test_predictor_async.py | 63 ++++++++++++++++++++++++++++++ 2 files changed, 91 insertions(+), 2 deletions(-) diff --git a/src/sagemaker/predictor_async.py b/src/sagemaker/predictor_async.py index 15b5d454d3..5eefa92dd9 100644 --- a/src/sagemaker/predictor_async.py +++ b/src/sagemaker/predictor_async.py @@ -57,14 +57,40 @@ def __init__( else: self.s3_client = self.sagemaker_session.s3_client - self.serializer = predictor.serializer - self.deserializer = predictor.deserializer self.name = name self._endpoint_config_name = None self._model_names = None self._context = None self._input_path = None + @property + def serializer(self): + """The serializer used to encode request data uploaded to Amazon S3. + + Reads from and writes to the wrapped ``Predictor``, so it always stays in + sync with the serializer the underlying predictor uses. + """ + return self.predictor.serializer + + @serializer.setter + def serializer(self, serializer): + """Set the serializer on the wrapped ``Predictor``.""" + self.predictor.serializer = serializer + + @property + def deserializer(self): + """The deserializer used to decode the async inference result. + + Reads from and writes to the wrapped ``Predictor``, so the ``Accept`` header + and the decoding of the Amazon S3 output both honor the configured value. + """ + return self.predictor.deserializer + + @deserializer.setter + def deserializer(self, deserializer): + """Set the deserializer on the wrapped ``Predictor``.""" + self.predictor.deserializer = deserializer + def predict( self, data=None, diff --git a/tests/unit/test_predictor_async.py b/tests/unit/test_predictor_async.py index c9f12ff023..629d641475 100644 --- a/tests/unit/test_predictor_async.py +++ b/tests/unit/test_predictor_async.py @@ -12,11 +12,15 @@ # language governing permissions and limitations under the License. from __future__ import absolute_import +import io + import pytest from mock import Mock from sagemaker.async_inference.waiter_config import WaiterConfig +from sagemaker.deserializers import JSONDeserializer from sagemaker.predictor import Predictor from sagemaker.predictor_async import AsyncPredictor +from sagemaker.serializers import JSONSerializer from sagemaker.exceptions import AsyncInferenceModelError, PollingTimeoutError ENDPOINT = "mxnet_endpoint" @@ -494,3 +498,62 @@ def test_list_monitors(): predictor_async.list_monitors() predictor.list_monitors.assert_called_with() + + +def _sagemaker_session_with_json_output(): + sagemaker_session = empty_sagemaker_session() + sagemaker_session.s3_client.get_waiter = Mock(name="get_waiter") + sagemaker_session.s3_client.get_object = Mock( + name="get_object", + return_value={ + "Body": io.BytesIO(b'{"result": [1, 2, 3]}'), + "ContentType": "application/json", + }, + ) + return sagemaker_session + + +def test_async_predictor_deserializer_override_is_used_for_accept_and_result(): + sagemaker_session = _sagemaker_session_with_json_output() + predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session)) + + predictor_async.deserializer = JSONDeserializer() + + assert isinstance(predictor_async.predictor.deserializer, JSONDeserializer) + + result = predictor_async.predict( + input_path=ASYNC_INPUT_LOCATION, waiter_config=DEFAULT_WAITER_CONFIG + ) + + _, kwargs = sagemaker_session.sagemaker_runtime_client.invoke_endpoint_async.call_args + assert kwargs["Accept"] == "application/json" + assert result == {"result": [1, 2, 3]} + + +def test_async_predictor_serializer_override_is_used_for_upload(): + sagemaker_session = empty_sagemaker_session() + predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session)) + predictor_async.name = ASYNC_PREDICTOR + + predictor_async.serializer = JSONSerializer() + + assert isinstance(predictor_async.predictor.serializer, JSONSerializer) + + predictor_async.predict_async(data={"hi": "there"}) + + _, kwargs = sagemaker_session.s3_client.put_object.call_args + assert kwargs["Body"] == '{"hi": "there"}' + assert kwargs["ContentType"] == "application/json" + + +def test_async_predictor_reflects_serializers_set_on_wrapped_predictor(): + predictor = Predictor(ENDPOINT, empty_sagemaker_session()) + predictor_async = AsyncPredictor(predictor) + + serializer = JSONSerializer() + deserializer = JSONDeserializer() + predictor.serializer = serializer + predictor.deserializer = deserializer + + assert predictor_async.serializer is serializer + assert predictor_async.deserializer is deserializer