Make command reference publication race-safe
This commit is contained in:
parent
efe4a443b7
commit
3bf3836ac3
2 changed files with 361 additions and 24 deletions
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue