Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 9 additions & 3 deletions sagemaker-core/src/sagemaker/core/processing.py
Original file line number Diff line number Diff line change
Expand Up @@ -1289,7 +1289,9 @@ def _package_code(
raise ValueError(f"source_dir does not exist: {source_dir}")

# Create tar.gz with source_dir contents + dependencies
with tempfile.NamedTemporaryFile(suffix=".tar.gz", delete=False) as tmp:
tmp = tempfile.NamedTemporaryFile(suffix=".tar.gz", delete=False)
tmp.close()
try:
with tarfile.open(tmp.name, "w:gz") as tar:
# Add all files from source_dir
for item in os.listdir(source_dir):
Expand All @@ -1311,15 +1313,19 @@ def _package_code(
"sourcedir.tar.gz",
)

with open(tmp.name, "rb") as tar_file:
tar_bytes = tar_file.read()

s3.S3Uploader.upload_string_as_file_body(
body=open(tmp.name, "rb").read(),
body=tar_bytes,
desired_s3_uri=s3_uri,
kms_key=kms_key,
sagemaker_session=self.sagemaker_session,
)

os.unlink(tmp.name)
return s3_uri
finally:
os.unlink(tmp.name)

@_telemetry_emitter(feature=Feature.PROCESSING, func_name="FrameworkProcessor.run")
@runnable_by_pipeline
Expand Down
53 changes: 53 additions & 0 deletions sagemaker-core/tests/unit/test_processing.py
Original file line number Diff line number Diff line change
Expand Up @@ -1290,6 +1290,59 @@ def test_package_code_with_code_location_trailing_slash(self, mock_session):
assert result.startswith("s3://my-custom-bucket/my-prefix")
assert "sourcedir.tar.gz" in result

def test_package_code_closes_temp_handle_before_unlink(self, mock_session):
"""Temp tar.gz must be closed before os.unlink (issue #5873).

On Windows os.unlink raises PermissionError (WinError 32) if any
handle to the file is still open. We track every open handle on the
temp path and assert none remain open when os.unlink is called.
"""
processor = FrameworkProcessor(
role="arn:aws:iam::123456789012:role/SageMakerRole",
image_uri="test-image:latest",
instance_count=1,
instance_type="ml.m5.xlarge",
sagemaker_session=mock_session,
)

real_open = open
open_handles = {}

def tracking_open(file, mode="r", *args, **kwargs):
handle = real_open(file, mode, *args, **kwargs)
if isinstance(file, str) and file.endswith(".tar.gz"):
open_handles[handle] = file
return handle

real_unlink = os.unlink
observed = {}

def checking_unlink(path, *args, **kwargs):
if isinstance(path, str) and path.endswith(".tar.gz"):
observed["still_open"] = [
p for h, p in open_handles.items() if p == path and not h.closed
]
return real_unlink(path, *args, **kwargs)

with tempfile.TemporaryDirectory() as tmpdir:
entry_point = os.path.join(tmpdir, "train.py")
with real_open(entry_point, "w") as f:
f.write("print('training')")

with patch("builtins.open", side_effect=tracking_open):
with patch("sagemaker.core.processing.os.unlink", side_effect=checking_unlink):
processor._package_code(
entry_point=entry_point,
source_dir=tmpdir,
requirements=None,
job_name="test-job",
kms_key=None,
)

assert (
observed.get("still_open") == []
), "temp tar.gz handle was still open when os.unlink was called"


class TestFrameworkProcessorRun:
def test_run_with_s3_code(self, mock_session):
Expand Down
Loading