1
0
Fork 0
Code Issues Pull requests Projects Releases 2 Packages Wiki Activity Actions Pages

Make command reference publication race-safe

This commit is contained in:
Andraxion 2026-07-29 15:00:28 -04:00
parent efe4a443b7
commit 3bf3836ac3
2 changed files with 361 additions and 24 deletions

View file

@ -6,10 +6,13 @@ import hashlib
import io
import os
import tempfile
import threading
import unittest
from pathlib import Path
from unittest import mock
from docforge.mcp_server import ALL_TOOLS, APPLICATION_TOOLS
from tools import generate_command_reference as command_reference_tool
from tools.generate_command_reference import (
EXPECTED_CLI_ROWS,
EXPECTED_MCP_ROWS,
@ -183,6 +186,173 @@ class CommandReferenceToolTests(unittest.TestCase):
check=False,
)
def test_existing_target_race_is_restored_without_overwrite(self) -> None:
with tempfile.TemporaryDirectory() as directory:
repository = Path(directory).resolve()
generated = repository / "generated"
generated.mkdir()
output = Path("generated/reference.md")
target = repository / output
target.write_bytes(b"initial\n")
original_exchange = command_reference_tool._rename_exchange
injected = False
def exchange_after_race(directory_fd: int, first: str, second: str) -> None:
nonlocal injected
if not injected:
injected = True
target.write_bytes(b"concurrent\n")
original_exchange(directory_fd, first, second)
with (
mock.patch.object(
command_reference_tool,
"_rename_exchange",
side_effect=exchange_after_race,
),
self.assertRaisesRegex(
CommandReferenceToolError,
"changed during atomic publication",
),
):
publish_or_check_command_reference(
b"generated\n",
repository_root=repository,
project_root=ALPHA,
output=output,
check=False,
)
self.assertEqual(b"concurrent\n", target.read_bytes())
self.assertEqual(
[],
list(generated.glob(".reference.md.docforge-command-reference-*")),
)
def test_missing_target_race_is_never_clobbered(self) -> None:
with tempfile.TemporaryDirectory() as directory:
repository = Path(directory).resolve()
generated = repository / "generated"
generated.mkdir()
output = Path("generated/reference.md")
target = repository / output
original_link = command_reference_tool._link_no_replace
injected = False
def link_after_race(directory_fd: int, source: str, destination: str) -> None:
nonlocal injected
if not injected:
injected = True
target.write_bytes(b"concurrent\n")
original_link(directory_fd, source, destination)
with (
mock.patch.object(
command_reference_tool,
"_link_no_replace",
side_effect=link_after_race,
),
self.assertRaisesRegex(CommandReferenceToolError, "appeared"),
):
publish_or_check_command_reference(
b"generated\n",
repository_root=repository,
project_root=ALPHA,
output=output,
check=False,
)
self.assertEqual(b"concurrent\n", target.read_bytes())
self.assertEqual(
[],
list(generated.glob(".reference.md.docforge-command-reference-*")),
)
def test_concurrent_generators_are_serialized_around_inspection_and_write(self) -> None:
with tempfile.TemporaryDirectory() as directory:
repository = Path(directory).resolve()
(repository / "generated").mkdir()
output = Path("generated/reference.md")
target = repository / output
target.write_bytes(b"initial\n")
first_entered = threading.Event()
release_first = threading.Event()
second_started = threading.Event()
second_entered = threading.Event()
results: list[str] = []
errors: list[BaseException] = []
original_replace = command_reference_tool._atomic_replace
invocation_count = 0
invocation_lock = threading.Lock()
def blocking_replace(
directory_fd: int,
name: str,
content: bytes,
*,
expected: object,
expected_content: bytes | None,
) -> None:
nonlocal invocation_count
with invocation_lock:
invocation_count += 1
invocation = invocation_count
if invocation == 1:
first_entered.set()
if not release_first.wait(timeout=5):
raise AssertionError("Timed out waiting to release first generator")
else:
second_entered.set()
original_replace(
directory_fd,
name,
content,
expected=expected, # type: ignore[arg-type]
expected_content=expected_content,
)
def publish(content: bytes, *, started: threading.Event | None = None) -> None:
if started is not None:
started.set()
try:
results.append(
publish_or_check_command_reference(
content,
repository_root=repository,
project_root=ALPHA,
output=output,
check=False,
)
)
except BaseException as error:
errors.append(error)
with mock.patch.object(
command_reference_tool,
"_atomic_replace",
side_effect=blocking_replace,
):
first = threading.Thread(target=publish, args=(b"first\n",))
second = threading.Thread(
target=publish,
args=(b"second\n",),
kwargs={"started": second_started},
)
first.start()
self.assertTrue(first_entered.wait(timeout=5))
second.start()
self.assertTrue(second_started.wait(timeout=5))
self.assertFalse(second_entered.wait(timeout=0.1))
release_first.set()
first.join(timeout=5)
second.join(timeout=5)
self.assertFalse(first.is_alive())
self.assertFalse(second.is_alive())
self.assertEqual([], errors)
self.assertEqual(["written", "written"], results)
self.assertEqual(b"second\n", target.read_bytes())
def test_command_runs_with_explicit_repository_and_project_roots(self) -> None:
with tempfile.TemporaryDirectory() as directory:
repository = Path(directory).resolve()