Skip to content
Merged
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
149 changes: 149 additions & 0 deletions tests/host_weight_runtime/test_filesystem_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from dataclasses import replace
from multiprocessing.process import BaseProcess
from pathlib import Path
from typing import IO, Any

import pytest
import torch
Expand Down Expand Up @@ -408,6 +409,154 @@ def _make_store(
return FilesystemHostWeightStore(domain, capacity or CapacityPolicy(), integrity or IntegrityPolicy())


class _FaultyTemporaryFile:
"""NamedTemporaryFile proxy that injects a write or close failure."""

def __init__(
self,
handle: IO[bytes],
*,
write_error: BaseException | None = None,
close_error: BaseException | None = None,
) -> None:
self._handle = handle
self._write_error = write_error
self._close_error = close_error

def __getattr__(self, name: str) -> object:
return getattr(self._handle, name)

def __enter__(self) -> _FaultyTemporaryFile:
return self

def __exit__(self, *exc_info: object) -> None:
self._handle.close()
if self._close_error is not None:
raise self._close_error

def write(self, data: bytes) -> int:
if self._write_error is not None:
raise self._write_error
return self._handle.write(data)


def _inject_temporary_file_failure(
monkeypatch: pytest.MonkeyPatch,
*,
write_error: BaseException | None = None,
close_error: BaseException | None = None,
) -> None:
original = tempfile.NamedTemporaryFile

def faulty(*args: Any, **kwargs: Any) -> _FaultyTemporaryFile:
handle: IO[bytes] = original(*args, **kwargs)
return _FaultyTemporaryFile(handle, write_error=write_error, close_error=close_error)

monkeypatch.setattr(tempfile, "NamedTemporaryFile", faulty)


def test_atomic_json_replace_failure_removes_temporary_file(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
path = tmp_path / "domain.json"
path.write_bytes(b"existing metadata")
replace_error = OSError(errno.EIO, "injected replace failure")

def fail_replace(source: os.PathLike[str], destination: os.PathLike[str]) -> None:
raise replace_error

monkeypatch.setattr(os, "replace", fail_replace)

with pytest.raises(OSError) as error:
filesystem_store_module._write_atomic_json(path, {"schema_version": 1})

assert error.value is replace_error
assert path.read_bytes() == b"existing metadata"
assert not list(tmp_path.glob(".domain.json.*.tmp"))


def test_atomic_json_cleanup_failure_preserves_replace_error(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
path = tmp_path / "domain.json"
replace_error = OSError(errno.EIO, "injected replace failure")
original_unlink = Path.unlink
cleanup_attempted = False

def fail_replace(source: os.PathLike[str], destination: os.PathLike[str]) -> None:
raise replace_error

def fail_temporary_unlink(target: Path, missing_ok: bool = False) -> None:
nonlocal cleanup_attempted
if target.parent == tmp_path and target.name.startswith(".domain.json."):
cleanup_attempted = True
raise PermissionError(errno.EACCES, "injected cleanup failure")
original_unlink(target, missing_ok=missing_ok)

monkeypatch.setattr(os, "replace", fail_replace)
monkeypatch.setattr(Path, "unlink", fail_temporary_unlink)

with pytest.raises(OSError) as error:
filesystem_store_module._write_atomic_json(path, {"schema_version": 1})

assert cleanup_attempted
assert error.value is replace_error
assert len(list(tmp_path.glob(".domain.json.*.tmp"))) == 1


@pytest.mark.parametrize("close_fails", [False, True])
def test_atomic_json_write_failure_preserves_error_when_cleanup_fails(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
close_fails: bool,
) -> None:
path = tmp_path / "domain.json"
write_error = OSError(errno.EIO, "injected write failure")
original_unlink = Path.unlink
cleanup_attempted = False

def fail_temporary_unlink(target: Path, missing_ok: bool = False) -> None:
nonlocal cleanup_attempted
if target.parent == tmp_path and target.name.startswith(".domain.json."):
cleanup_attempted = True
raise PermissionError(errno.EACCES, "injected cleanup failure")
original_unlink(target, missing_ok=missing_ok)

_inject_temporary_file_failure(
monkeypatch,
write_error=write_error,
close_error=OSError(errno.EIO, "injected close failure") if close_fails else None,
)
monkeypatch.setattr(Path, "unlink", fail_temporary_unlink)

with pytest.raises(OSError) as error:
filesystem_store_module._write_atomic_json(path, {"schema_version": 1})

assert cleanup_attempted
assert error.value is write_error
assert not path.exists()
assert len(list(tmp_path.glob(".domain.json.*.tmp"))) == 1


def test_atomic_json_close_failure_removes_temporary_file(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
path = tmp_path / "domain.json"
close_error = OSError(errno.EIO, "injected close failure")

_inject_temporary_file_failure(monkeypatch, close_error=close_error)

with pytest.raises(OSError) as error:
filesystem_store_module._write_atomic_json(path, {"schema_version": 1})

assert error.value is close_error
assert not path.exists()
assert not list(tmp_path.glob(".domain.json.*.tmp"))


def _publish_test_artifact(
store: FilesystemHostWeightStore,
identity: WeightArtifactIdentity | None = None,
Expand Down
35 changes: 19 additions & 16 deletions vllm_omni/host_weight_runtime/filesystem/store.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
import os
import shutil
import stat
import tempfile
import time
import uuid
from collections import defaultdict
Expand Down Expand Up @@ -298,25 +299,27 @@ def _move_artifact_locked(source: Path, destination: Path) -> None:

def _write_atomic_json(path: Path, value: object, *, mode: int = 0o444) -> None:
data = canonical_json(value)
temporary = path.parent / f".{path.name}.{os.getpid()}.{uuid.uuid4().hex}.tmp"
fd = os.open(temporary, os.O_CREAT | os.O_EXCL | os.O_WRONLY | os.O_CLOEXEC, 0o600)
handle = tempfile.NamedTemporaryFile(dir=path.parent, prefix=f".{path.name}.", suffix=".tmp", delete=False)
temporary = Path(handle.name)
write_error: BaseException | None = None
try:
offset = 0
while offset < len(data):
count = os.write(fd, data[offset:])
if count <= 0:
raise OSError(errno.EIO, f"short write for {temporary}")
offset += count
os.fsync(fd)
os.fchmod(fd, mode)
os.fsync(fd)
with handle:
try:
handle.write(data)
handle.flush()
os.fsync(handle.fileno())
os.fchmod(handle.fileno(), mode)
os.fsync(handle.fileno())
except BaseException as error:
write_error = error
raise
os.replace(temporary, path)
except BaseException:
os.close(fd)
temporary.unlink(missing_ok=True)
with contextlib.suppress(OSError):
temporary.unlink(missing_ok=True)
if write_error is not None:
raise write_error
raise
else:
os.close(fd)
os.replace(temporary, path)
_fsync_directory(path.parent)


Expand Down
Loading