diff --git a/path/__init__.py b/path/__init__.py index 33dd979..64f8c56 100644 --- a/path/__init__.py +++ b/path/__init__.py @@ -1842,15 +1842,20 @@ def only_newer(copy_func: _CopyFn) -> _CopyFn: """ Wrap a copy function (like shutil.copy2) to return the dst if it's newer than the source. + + If dst is a directory, compare against the file with the source's name + inside that directory. """ @functools.wraps(copy_func) def wrapper(src: str, dst: str): src_p = Path(src) dst_p = Path(dst) + if dst_p.is_dir(): + dst_p = dst_p / src_p.name is_newer_dst = dst_p.exists() and dst_p.getmtime() >= src_p.getmtime() if is_newer_dst: - return dst + return dst_p return copy_func(src, dst) return wrapper diff --git a/tests/test_path.py b/tests/test_path.py index b424bb6..cc0c4e5 100644 --- a/tests/test_path.py +++ b/tests/test_path.py @@ -1398,3 +1398,36 @@ def test_ignore(self): def test_invalid_handler(self): with pytest.raises(ValueError): path.Handlers._resolve('raise') + + +@pytest.mark.parametrize('target_mtime', [None, 1000, 3000]) +def test_only_newer_directory(tmp_path, target_mtime): + source = Path(tmp_path) / 'source.txt' + source.write_text('source') + source.utime((2000, 2000)) + directory = (Path(tmp_path) / 'destination').mkdir() + target = directory / source.name + if target_mtime is not None: + target.write_text('destination') + target.utime((target_mtime, target_mtime)) + directory.utime((4000, 4000)) + + result = path.only_newer(shutil.copy2)(source, directory) + + assert result == target + expected = 'destination' if target_mtime == 3000 else 'source' + assert target.read_text() == expected + + +def test_only_newer_directory_with_newer_file(tmp_path): + source = Path(tmp_path) / 'source.txt' + source.write_text('source') + source.utime((2000, 2000)) + directory = (Path(tmp_path) / 'destination').mkdir() + target = directory / source.name + target.write_text('destination') + target.utime((3000, 3000)) + directory.utime((1000, 1000)) + + assert path.only_newer(shutil.copy2)(source, directory) == target + assert target.read_text() == 'destination'