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
104 changes: 104 additions & 0 deletions autotest/test_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,13 +9,16 @@
import os
import subprocess
import sys
import threading
import zipfile
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path

import pooch
import pytest
from flaky import flaky

import modflow_devtools.models as models_module
from modflow_devtools.models import (
_DEFAULT_BASE_URL,
_DEFAULT_CACHE,
Expand Down Expand Up @@ -436,6 +439,107 @@ def test_registry_to_pooch_format(self, synced_registry):
assert len(pooch_registry) == len(synced_registry.files)


class TestConcurrentFetch:
"""Test that fetching model files is safe under concurrent workers."""

MODEL = "mf6/model"
FILES = ("mf6/model/a.dat", "mf6/model/b.dat")

@pytest.fixture
def registry(self, tmp_path, monkeypatch):
"""A PoochRegistry on a cold tmp cache with a fake, offline downloader."""
content = b"data"
digest = hashlib.sha256(content).hexdigest()

def fake_downloader(url, output_file, pooch, check_only=False):
Path(output_file).write_bytes(content)

monkeypatch.setattr(PoochRegistry, "_load", lambda self: None)
monkeypatch.setattr(pooch.core, "choose_downloader", lambda *a, **kw: fake_downloader)

registry = PoochRegistry(path=tmp_path / "cache", base_url="https://example.invalid/")
registry.models[self.MODEL] = list(self.FILES)
registry.pooch.registry = dict.fromkeys(self.FILES, digest)
registry.pooch.urls = {f: f"https://example.invalid/{f}" for f in self.FILES}
return registry

def test_fetch_files_concurrent_new_subdirectory(self, registry, monkeypatch):
"""
Workers fetching into the same new subdirectory must not race on makedirs.

Pooch checks whether a file's parent directory exists, then calls
os.makedirs without exist_ok. The wrapper below holds every worker at
makedirs until all have passed the existence check, forcing the race.
Without pre-creating the directory this fails every time; with it,
makedirs is never reached, so the barrier is never used.
"""
n = 8
barrier = threading.Barrier(n)
real_makedirs = os.makedirs

def racing_makedirs(*args, **kwargs):
try:
barrier.wait(timeout=10)
except threading.BrokenBarrierError:
pass
return real_makedirs(*args, **kwargs)

monkeypatch.setattr(pooch.core.os, "makedirs", racing_makedirs)

fetch = registry._fetcher(self.MODEL, list(self.FILES))
with ThreadPoolExecutor(max_workers=n) as pool:
results = list(pool.map(lambda _: fetch(), range(n)))

expected = [registry.pooch.abspath / f for f in self.FILES]
for paths in results:
assert paths == expected
assert all(p.is_file() for p in expected)

def test_fetch_zip_lock_in_cache_dir(self, registry, tmp_path, monkeypatch):
"""The zip lock file belongs in the cache next to the zip, not in the cwd."""
buf = io.BytesIO()
with zipfile.ZipFile(buf, "w") as zf:
for f in self.FILES:
zf.writestr(f, "data")
zip_bytes = buf.getvalue()

def zip_downloader(url, output_file, pooch, check_only=False):
Path(output_file).write_bytes(zip_bytes)

monkeypatch.setattr(pooch.core, "choose_downloader", lambda *a, **kw: zip_downloader)

# an empty hash for every model file selects the zip fetcher
registry.pooch.registry = dict.fromkeys(self.FILES, "")
registry.pooch.registry[_DEFAULT_ZIP_NAME] = hashlib.sha256(zip_bytes).hexdigest()
registry.pooch.urls = {_DEFAULT_ZIP_NAME: f"https://example.invalid/{_DEFAULT_ZIP_NAME}"}

lock_paths = []

class RecordingLock:
def __init__(self, path, *args, **kwargs):
lock_paths.append(Path(path))

def __enter__(self):
return self

def __exit__(self, *args):
return False

monkeypatch.setattr(models_module, "FileLock", RecordingLock)

cwd = tmp_path / "cwd"
cwd.mkdir()
monkeypatch.chdir(cwd)

fetch = registry._fetcher(self.MODEL, list(self.FILES))
paths = fetch()

assert len(paths) == len(self.FILES)
assert lock_paths == [registry.pooch.abspath / f"{_DEFAULT_ZIP_NAME}.lock"]
assert lock_paths[0].parent.is_dir()
assert not any(cwd.iterdir())


@pytest.mark.xdist_group("registry_cache")
class TestCLI:
"""Test CLI commands."""
Expand Down
13 changes: 12 additions & 1 deletion modflow_devtools/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -1036,10 +1036,21 @@ def __init__(

def _fetcher(self, model_name, file_names) -> Callable:
def _fetch_files():
# Pooch creates a file's parent directory with a bare
# check-then-makedirs, which raises FileExistsError if
# concurrent workers (e.g. pytest-xdist) fetch into the
# same new subdirectory. Creating it first, tolerantly,
# means pooch's existence check passes and it skips it.
for fname in file_names:
(self.pooch.abspath / fname).parent.mkdir(parents=True, exist_ok=True)
return [Path(self.pooch.fetch(fname)) for fname in file_names]

def _fetch_zip(zip_name):
with FileLock(f"{zip_name}.lock"):
# Keep the lock in the cache next to the zip it protects, so it
# excludes processes regardless of cwd and doesn't litter it.
lock_path = self.pooch.abspath / f"{zip_name}.lock"
lock_path.parent.mkdir(parents=True, exist_ok=True)
with FileLock(str(lock_path)):
return [
Path(f)
for f in self.pooch.fetch(
Expand Down
Loading