diff --git a/sagemaker-serve/src/sagemaker/serve/model_server/multi_model_server/prepare.py b/sagemaker-serve/src/sagemaker/serve/model_server/multi_model_server/prepare.py index 3b347ee65c..4821762517 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_server/multi_model_server/prepare.py +++ b/sagemaker-serve/src/sagemaker/serve/model_server/multi_model_server/prepare.py @@ -123,3 +123,5 @@ def prepare_for_mms( hash_value = compute_hash(buffer=buffer) with open(str(code_dir.joinpath("metadata.json")), "wb") as metadata: metadata.write(_MetaData(hash_value).to_json()) + + return hash_value diff --git a/sagemaker-serve/src/sagemaker/serve/model_server/smd/prepare.py b/sagemaker-serve/src/sagemaker/serve/model_server/smd/prepare.py index f29b8ebcbd..81e81fc926 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_server/smd/prepare.py +++ b/sagemaker-serve/src/sagemaker/serve/model_server/smd/prepare.py @@ -68,3 +68,5 @@ def prepare_for_smd( hash_value = compute_hash(buffer=buffer) with open(str(code_dir.joinpath("metadata.json")), "wb") as metadata: metadata.write(_MetaData(hash_value).to_json()) + + return hash_value diff --git a/sagemaker-serve/src/sagemaker/serve/model_server/tensorflow_serving/prepare.py b/sagemaker-serve/src/sagemaker/serve/model_server/tensorflow_serving/prepare.py index d56d0ec7bd..5345331812 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_server/tensorflow_serving/prepare.py +++ b/sagemaker-serve/src/sagemaker/serve/model_server/tensorflow_serving/prepare.py @@ -61,3 +61,5 @@ def prepare_for_tf_serving( hash_value = compute_hash(buffer=buffer) with open(str(code_dir.joinpath("metadata.json")), "wb") as metadata: metadata.write(_MetaData(hash_value).to_json()) + + return hash_value diff --git a/sagemaker-serve/src/sagemaker/serve/model_server/torchserve/prepare.py b/sagemaker-serve/src/sagemaker/serve/model_server/torchserve/prepare.py index ad053d25c9..55b9d4cc36 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_server/torchserve/prepare.py +++ b/sagemaker-serve/src/sagemaker/serve/model_server/torchserve/prepare.py @@ -73,3 +73,5 @@ def prepare_for_torchserve( hash_value = compute_hash(buffer=buffer) with open(str(code_dir.joinpath("metadata.json")), "wb") as metadata: metadata.write(_MetaData(hash_value).to_json()) + + return hash_value diff --git a/sagemaker-serve/tests/unit/model_server/test_smd_prepare.py b/sagemaker-serve/tests/unit/model_server/test_smd_prepare.py index aa21763180..a09c9957e4 100644 --- a/sagemaker-serve/tests/unit/model_server/test_smd_prepare.py +++ b/sagemaker-serve/tests/unit/model_server/test_smd_prepare.py @@ -44,6 +44,7 @@ def test_prepare_for_smd_with_inference_spec(self, mock_copy, mock_capture, mock ) mock_inference_spec.prepare.assert_called_once_with(str(model_path)) + self.assertEqual(secret_key, "test-hash") @patch("os.rename") @patch("sagemaker.serve.model_server.smd.prepare.compute_hash") @@ -76,6 +77,7 @@ def test_prepare_for_smd_with_custom_orchestrator( # Verify custom_execution_inference.py was copied and renamed mock_rename.assert_called_once() + self.assertEqual(secret_key, "test-hash") @patch("sagemaker.serve.model_server.smd.prepare.compute_hash") @patch("sagemaker.serve.model_server.smd.prepare.capture_dependencies") @@ -115,6 +117,31 @@ def test_prepare_for_smd_invalid_dir(self): prepare_for_smd(model_path=str(file_path), shared_libs=[], dependencies={}) self.assertIn("not a valid directory", str(context.exception)) + @patch("sagemaker.serve.model_server.smd.prepare.capture_dependencies") + @patch("shutil.copy2") + def test_prepare_for_smd_returns_hash_value(self, mock_copy, mock_capture): + """Test prepare_for_smd returns a valid SHA-256 hash string.""" + from sagemaker.serve.model_server.smd.prepare import prepare_for_smd + + model_path = Path(self.temp_dir) / "model" + code_dir = model_path / "code" + code_dir.mkdir(parents=True) + + # Create a real serve.pkl file + serve_pkl = code_dir / "serve.pkl" + serve_pkl.write_bytes(b"test pickle data") + + result = prepare_for_smd( + model_path=str(model_path), shared_libs=[], dependencies={} + ) + + # Verify the return value is a valid 64-character hex string (SHA-256) + self.assertIsNotNone(result) + self.assertIsInstance(result, str) + self.assertEqual(len(result), 64) + # Verify it's a valid hex string + int(result, 16) + if __name__ == "__main__": unittest.main()