From 6dba585356a7dbe6fc1297a5a28e2402e263daa6 Mon Sep 17 00:00:00 2001 From: Lucas Jia Date: Fri, 25 Sep 2026 14:19:37 -0700 Subject: [PATCH] fix: Default AsyncPredictor upload prefix to endpoint name AsyncPredictor accepts name=None by default, but predict() and predict_async() called name_from_base(self.name) when uploading input data without an input_path, which raised "TypeError: 'NoneType' object is not subscriptable". Fall back to the wrapped predictor's endpoint name when no name is given, and document the name argument. An explicitly provided name is still used as the S3 key prefix. Fixes #3210 Fixes #4774 --- src/sagemaker/predictor_async.py | 8 +++++++- tests/unit/test_predictor_async.py | 25 +++++++++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/src/sagemaker/predictor_async.py b/src/sagemaker/predictor_async.py index 15b5d454d3..9eebdf781e 100644 --- a/src/sagemaker/predictor_async.py +++ b/src/sagemaker/predictor_async.py @@ -40,6 +40,9 @@ def __init__( predictor (sagemaker.predictor.Predictor): General ``Predictor`` object has useful methods and variables. ``AsyncPredictor`` stands on top of it with capability for async inference. + name (str): Optional. Name used as the prefix of the Amazon S3 key when + input data is uploaded for async inference. If not specified, the + endpoint name is used. (Default: None) """ warn_v2_deprecation( feature="AsyncPredictor", @@ -174,10 +177,13 @@ def _upload_data_to_s3( my_uuid = str(uuid.uuid4()) timestamp = sagemaker_timestamp() bucket = self.sagemaker_session.default_bucket() + # ``name`` is optional; fall back to the endpoint name so the default + # ``AsyncPredictor(predictor)`` construction can upload input data. + base_name = self.name or self.endpoint_name key = s3.s3_path_join( self.sagemaker_session.default_bucket_prefix, "async-endpoint-inputs", - name_from_base(self.name, short=True), + name_from_base(base_name, short=True), "{}-{}".format(timestamp, my_uuid), ) diff --git a/tests/unit/test_predictor_async.py b/tests/unit/test_predictor_async.py index c9f12ff023..eb229a0cc9 100644 --- a/tests/unit/test_predictor_async.py +++ b/tests/unit/test_predictor_async.py @@ -170,6 +170,31 @@ def test_async_predict_call_with_data(): assert result.output_path == ASYNC_OUTPUT_LOCATION +def test_async_predict_call_with_data_and_no_name(): + # Regression test for https://github.com/aws/sagemaker-python-sdk/issues/3210 + sagemaker_session = empty_sagemaker_session() + predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session)) + assert predictor_async.name is None + + result = predictor_async.predict_async(data=DUMMY_DATA) + + _, put_kwargs = sagemaker_session.s3_client.put_object.call_args + assert put_kwargs["Bucket"] == BUCKET_NAME + assert put_kwargs["Key"].startswith("async-endpoint-inputs/{}-".format(ENDPOINT)) + assert predictor_async._input_path == "s3://{}/{}".format(BUCKET_NAME, put_kwargs["Key"]) + assert result.output_path == ASYNC_OUTPUT_LOCATION + + +def test_async_predict_call_with_data_and_name_uses_name(): + sagemaker_session = empty_sagemaker_session() + predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session), name=ASYNC_PREDICTOR) + + predictor_async.predict_async(data=DUMMY_DATA) + + _, put_kwargs = sagemaker_session.s3_client.put_object.call_args + assert put_kwargs["Key"].startswith("async-endpoint-inputs/{}-".format(ASYNC_PREDICTOR)) + + def test_async_predict_call_with_data_and_input_path(): sagemaker_session = empty_sagemaker_session() predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session))