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))