diff --git a/git/index/base.py b/git/index/base.py index c9d99840d..4ecd0c5b1 100644 --- a/git/index/base.py +++ b/git/index/base.py @@ -1303,7 +1303,7 @@ def _flush_stdin_and_wait(self, proc: "Popen[bytes]", ignore_stdout: bool = Fals @default_index def checkout( self, - paths: Union[None, Iterable[PathLike]] = None, + paths: Union[None, PathLike, Iterable[PathLike]] = None, force: bool = False, fprogress: Callable = lambda *args: None, allow_unsafe_options: bool = False, @@ -1442,7 +1442,7 @@ def handle_stderr(proc: "Popen[bytes]", iter_checked_out_files: Iterable[PathLik handle_stderr(proc, rval_iter) return rval_iter else: - if isinstance(paths, str): + if isinstance(paths, (str, os.PathLike)): paths = [paths] # Make sure we have our entries loaded before we start checkout_index, which diff --git a/test/test_index.py b/test/test_index.py index 8c33fe29d..30cb54565 100644 --- a/test/test_index.py +++ b/test/test_index.py @@ -1731,6 +1731,32 @@ def test_index_file_v3_with_git_command(self, tmp_dir): assert "A file2.txt" in status_lines +class TestIndexCheckout: + @pytest.mark.parametrize("path_type", [str, Path, PathLikeMock]) + @pytest.mark.parametrize("absolute", [False, True]) + @pytest.mark.parametrize("directory", [False, True]) + @pytest.mark.parametrize("container", ["single", "list", "iterator"]) + def test_checkout_pathlike(self, tmp_path, path_type, absolute, directory, container): + with Repo.init(tmp_path) as repo: + nested = tmp_path / "nested" + nested.mkdir() + files = {"nested/first": b"first", "nested/second": b"second", "outside": b"outside"} + for name, data in files.items(): + (tmp_path / name).write_bytes(data) + repo.index.add(list(files)) + + path_name = "nested" if directory else "nested/first" + path = path_type(str(tmp_path / path_name) if absolute else path_name) + paths = path if container == "single" else [path] if container == "list" else iter([path]) + expected = {"nested/first", "nested/second"} if directory else {"nested/first"} + for name in expected: + (tmp_path / name).unlink() + + assert set(repo.index.checkout(paths)) == expected + for name, data in files.items(): + assert (tmp_path / name).read_bytes() == data + + class TestIndexUtils: @pytest.mark.parametrize("file_path_type", [str, Path]) def test_temporary_file_swap(self, tmp_path, file_path_type):