From 538d95cd0389b6b12ad40c3f22bf5a733338f86d Mon Sep 17 00:00:00 2001 From: Lucas Jia Date: Fri, 25 Sep 2026 13:22:47 -0700 Subject: [PATCH] fix: Keep AsyncPredictor serializers in sync with wrapped predictor Port the V2 fix for #3100 to sagemaker-serve. AsyncPredictor copied the wrapped predictor's serializer and deserializer at construction, so overriding AsyncPredictor.deserializer did not affect the Accept header or result decoding, which read from the wrapped predictor. Make serializer and deserializer properties that delegate to the wrapped predictor. Fixes #3100 --- .../src/sagemaker/serve/predictor_async.py | 30 +++++++- .../tests/unit/test_predictor_async.py | 77 +++++++++++++++++++ 2 files changed, 105 insertions(+), 2 deletions(-) diff --git a/sagemaker-serve/src/sagemaker/serve/predictor_async.py b/sagemaker-serve/src/sagemaker/serve/predictor_async.py index 31e94318d5..43130cd05b 100644 --- a/sagemaker-serve/src/sagemaker/serve/predictor_async.py +++ b/sagemaker-serve/src/sagemaker/serve/predictor_async.py @@ -52,14 +52,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/sagemaker-serve/tests/unit/test_predictor_async.py b/sagemaker-serve/tests/unit/test_predictor_async.py index ccd011a973..d4176ffc81 100644 --- a/sagemaker-serve/tests/unit/test_predictor_async.py +++ b/sagemaker-serve/tests/unit/test_predictor_async.py @@ -1,5 +1,9 @@ +import io import unittest from unittest.mock import Mock, patch + +from sagemaker.core.deserializers import JSONDeserializer +from sagemaker.core.serializers import JSONSerializer from sagemaker.serve.predictor_async import AsyncPredictor @@ -107,5 +111,78 @@ def test_disable_data_capture(self): self.mock_predictor.disable_data_capture.assert_called_once() +class _StubPredictor: + """Minimal predictor exposing the attributes ``AsyncPredictor`` relies on.""" + + def __init__(self, sagemaker_session, serializer, deserializer): + self.endpoint_name = "test-endpoint" + self.sagemaker_session = sagemaker_session + self.serializer = serializer + self.deserializer = deserializer + + @property + def accept(self): + return self.deserializer.ACCEPT + + def _handle_response(self, response): + return self.deserializer.deserialize(response["Body"], response["ContentType"]) + + +class TestAsyncPredictorSerializerOverrides(unittest.TestCase): + """Overrides set on AsyncPredictor must reach the request and the result (issue #3100).""" + + def setUp(self): + self.sagemaker_session = Mock() + self.sagemaker_session.default_bucket.return_value = "bucket" + self.sagemaker_session.default_bucket_prefix = None + self.sagemaker_session.sagemaker_runtime_client.invoke_endpoint_async.return_value = { + "OutputLocation": "s3://bucket/output", + } + self.sagemaker_session.s3_client.get_object.return_value = { + "Body": io.BytesIO(b'{"result": [1, 2, 3]}'), + "ContentType": "application/json", + } + default_deserializer = Mock() + default_deserializer.ACCEPT = ("*/*",) + self.predictor = _StubPredictor(self.sagemaker_session, Mock(), default_deserializer) + + def test_deserializer_override_is_used_for_accept_and_result(self): + async_predictor = AsyncPredictor(self.predictor) + + async_predictor.deserializer = JSONDeserializer() + + self.assertIsInstance(self.predictor.deserializer, JSONDeserializer) + + result = async_predictor.predict(input_path="s3://bucket/input") + + _, kwargs = self.sagemaker_session.sagemaker_runtime_client.invoke_endpoint_async.call_args + self.assertEqual(kwargs["Accept"], "application/json") + self.assertEqual(result, {"result": [1, 2, 3]}) + + def test_serializer_override_is_used_for_upload(self): + async_predictor = AsyncPredictor(self.predictor, name="test") + + async_predictor.serializer = JSONSerializer() + + self.assertIsInstance(self.predictor.serializer, JSONSerializer) + + async_predictor.predict_async(data={"hi": "there"}) + + _, kwargs = self.sagemaker_session.s3_client.put_object.call_args + self.assertEqual(kwargs["Body"], '{"hi": "there"}') + self.assertEqual(kwargs["ContentType"], "application/json") + + def test_reflects_serializers_set_on_wrapped_predictor(self): + async_predictor = AsyncPredictor(self.predictor) + + serializer = JSONSerializer() + deserializer = JSONDeserializer() + self.predictor.serializer = serializer + self.predictor.deserializer = deserializer + + self.assertIs(async_predictor.serializer, serializer) + self.assertIs(async_predictor.deserializer, deserializer) + + if __name__ == "__main__": unittest.main()