Enable cache when running with --diff (#1145) (#5499)

Co-authored-by: cobalt <61329810+cobaltt7@users.noreply.github.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
diff --git a/CHANGES.md b/CHANGES.md
index c5e4152..7e79d70 100644
--- a/CHANGES.md
+++ b/CHANGES.md
@@ -24,6 +24,9 @@
 
 <!-- Changes to how Black can be configured -->
 
+- Enable the cache when using `--diff` to skip unmodified files and record unmodified
+  files as well-formatted (#5499)
+
 ### Packaging
 
 <!-- Changes to how Black is packaged, such as dependency requirements -->
diff --git a/docs/usage_and_configuration/file_collection_and_discovery.md b/docs/usage_and_configuration/file_collection_and_discovery.md
index 330e1fe..a5e8b41 100644
--- a/docs/usage_and_configuration/file_collection_and_discovery.md
+++ b/docs/usage_and_configuration/file_collection_and_discovery.md
@@ -7,8 +7,8 @@
 
 ## Ignoring unmodified files
 
-_Black_ remembers files it has already formatted, unless the `--diff` flag is used or
-code is passed via standard input. This information is stored per-user. The exact
+_Black_ remembers files it has already formatted, unless code is passed via standard
+input or the `--no-cache` flag is used. This information is stored per-user. The exact
 location of the file depends on the _Black_ version and the system on which _Black_ is
 run. The file is non-portable. The standard location on common operating systems is:
 
diff --git a/src/black/__init__.py b/src/black/__init__.py
index 675ea87..847c4d2 100644
--- a/src/black/__init__.py
+++ b/src/black/__init__.py
@@ -1013,16 +1013,17 @@ def reformat_one(
                 cache = Cache.read(mode)
             else:
                 cache = Cache.read(mode, cache_dir)
-            if cache is not None and write_back not in (
-                WriteBack.DIFF,
-                WriteBack.COLOR_DIFF,
-            ):
+            if cache is not None:
                 if not cache.is_changed(src):
                     changed = Changed.CACHED
             if changed is not Changed.CACHED and format_file_in_place(
                 src, fast=fast, write_back=write_back, mode=mode, lines=lines
             ):
                 changed = Changed.YES
+            can_cache_unmodified = (
+                write_back in (WriteBack.CHECK, WriteBack.DIFF, WriteBack.COLOR_DIFF)
+                and changed is Changed.NO
+            )
             # Formatting only some lines doesn't make the whole file formatted, and
             # the cache key doesn't include the line ranges, so don't record it.
             if (
@@ -1030,7 +1031,7 @@ def reformat_one(
                 and not lines
                 and (
                     (write_back is WriteBack.YES and changed is not Changed.CACHED)
-                    or (write_back is WriteBack.CHECK and changed is Changed.NO)
+                    or can_cache_unmodified
                 )
             ):
                 cache.write([src])
diff --git a/src/black/concurrency.py b/src/black/concurrency.py
index 9953030..88bfda8 100644
--- a/src/black/concurrency.py
+++ b/src/black/concurrency.py
@@ -180,10 +180,7 @@ async def schedule_formatting(
         cache = Cache.read(mode)
     else:
         cache = Cache.read(mode, cache_dir)
-    if cache is not None and write_back not in (
-        WriteBack.DIFF,
-        WriteBack.COLOR_DIFF,
-    ):
+    if cache is not None:
         sources, cached = cache.filtered_cached(sources)
         for src in sorted(cached):
             report.done(src, Changed.CACHED)
@@ -237,11 +234,15 @@ async def schedule_formatting(
                 report.failed(src, exc)
             else:
                 changed = Changed.YES if task.result() else Changed.NO
-                # If the file was written back or was successfully checked as
-                # well-formatted, store this information in the cache.
-                if write_back is WriteBack.YES or (
-                    write_back is WriteBack.CHECK and changed is Changed.NO
-                ):
+                can_cache_unmodified = (
+                    write_back in (
+                        WriteBack.CHECK,
+                        WriteBack.DIFF,
+                        WriteBack.COLOR_DIFF,
+                    )
+                    and changed is Changed.NO
+                )
+                if write_back is WriteBack.YES or can_cache_unmodified:
                     sources_to_cache.append(src)
                 report.done(src, changed)
         if cancelled:
diff --git a/tests/test_black.py b/tests/test_black.py
index 439e6ea..b3b6f2f 100644
--- a/tests/test_black.py
+++ b/tests/test_black.py
@@ -2989,6 +2989,62 @@ def test_no_cache_when_writeback_diff(self, color: bool) -> None:
                 write_cache.assert_not_called()
 
     @pytest.mark.parametrize("color", [False, True], ids=["no-color", "with-color"])
+    def test_cache_written_when_writeback_diff_and_unmodified(
+        self, color: bool
+    ) -> None:
+        mode = DEFAULT_MODE
+        with cache_dir() as workspace:
+            src = (workspace / "test.py").resolve()
+            src.write_text('print("hello")\n', encoding="utf-8")
+            cmd = [str(src), "--diff"]
+            if color:
+                cmd.append("--color")
+            invokeBlack(cmd)
+            cache = black.Cache.read(mode, workspace)
+            assert not cache.is_changed(src)
+
+    @pytest.mark.parametrize("color", [False, True], ids=["no-color", "with-color"])
+    def test_cache_used_when_writeback_diff_already_cached(self, color: bool) -> None:
+        mode = DEFAULT_MODE
+        with cache_dir() as workspace:
+            src = (workspace / "test.py").resolve()
+            src.write_text("print('hello')", encoding="utf-8")
+            cache = black.Cache.read(mode, workspace)
+            cache.write([src])
+            cmd = ["--config", str(THIS_DIR / "empty.toml"), str(src), "--diff"]
+            if color:
+                cmd.append("--color")
+            result = BlackRunner().invoke(black.main, cmd)
+            assert result.exit_code == 0
+            assert "1 file would be left unchanged" in result.output
+            assert "@@" not in result.output
+
+    @pytest.mark.parametrize("color", [False, True], ids=["no-color", "with-color"])
+    @event_loop()
+    def test_cache_multiple_files_when_writeback_diff(self, color: bool) -> None:
+        mode = DEFAULT_MODE
+        with (
+            cache_dir() as workspace,
+            patch("concurrent.futures.ProcessPoolExecutor", new=ThreadPoolExecutor),
+        ):
+            one = (workspace / "one.py").resolve()
+            one.write_text('print("hello")\n', encoding="utf-8")
+            two = (workspace / "two.py").resolve()
+            two.write_text("print('world')", encoding="utf-8")
+            cache = black.Cache.read(mode, workspace)
+            cache.write([one])
+            cmd = ["--config", str(THIS_DIR / "empty.toml"), "--diff", str(workspace)]
+            if color:
+                cmd.append("--color")
+            result = BlackRunner().invoke(black.main, cmd)
+            assert result.exit_code == 0
+            expected = "1 file would be reformatted, 1 file would be left unchanged."
+            assert expected in result.output
+            cache = black.Cache.read(mode, workspace)
+            assert not cache.is_changed(one)
+            assert cache.is_changed(two)
+
+    @pytest.mark.parametrize("color", [False, True], ids=["no-color", "with-color"])
     @event_loop()
     def test_output_locking_when_writeback_diff(self, color: bool) -> None:
         with cache_dir() as workspace: