402 lines
15 KiB
Python
402 lines
15 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import contextlib
|
|
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,
|
|
CommandReferenceDrift,
|
|
CommandReferenceToolError,
|
|
collect_command_reference_rows,
|
|
generate_command_reference_bytes,
|
|
main,
|
|
publish_or_check_command_reference,
|
|
)
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
ALPHA = ROOT / "tests" / "fixtures" / "alpha"
|
|
|
|
|
|
def _tree_snapshot(root: Path) -> tuple[tuple[str, str, str], ...]:
|
|
records: list[tuple[str, str, str]] = []
|
|
for path in sorted(root.rglob("*")):
|
|
relative = path.relative_to(root).as_posix()
|
|
if path.is_symlink():
|
|
records.append((relative, "symlink", os.readlink(path)))
|
|
elif path.is_dir():
|
|
records.append((relative, "directory", ""))
|
|
else:
|
|
records.append(
|
|
(
|
|
relative,
|
|
"file",
|
|
hashlib.sha256(path.read_bytes()).hexdigest(),
|
|
)
|
|
)
|
|
return tuple(records)
|
|
|
|
|
|
class CommandReferenceToolTests(unittest.TestCase):
|
|
def test_generated_bytes_are_stable_and_never_mutate_the_fixture(self) -> None:
|
|
before = _tree_snapshot(ALPHA)
|
|
|
|
first = generate_command_reference_bytes(ALPHA)
|
|
middle = _tree_snapshot(ALPHA)
|
|
second = generate_command_reference_bytes(ALPHA)
|
|
|
|
self.assertEqual(first, second)
|
|
self.assertEqual(before, middle)
|
|
self.assertEqual(before, _tree_snapshot(ALPHA))
|
|
self.assertTrue(first.endswith(b"\n"))
|
|
self.assertEqual(1, first.count(b"## CLI commands"))
|
|
self.assertEqual(1, first.count(b"## MCP tools"))
|
|
|
|
def test_registration_collection_is_complete(self) -> None:
|
|
cli_rows, mcp_rows = asyncio.run(
|
|
collect_command_reference_rows(
|
|
ALPHA,
|
|
proposal_writer="alpha-editor",
|
|
canonical_applier="alpha-editor",
|
|
)
|
|
)
|
|
|
|
self.assertEqual(EXPECTED_CLI_ROWS, len(cli_rows))
|
|
self.assertEqual(EXPECTED_MCP_ROWS, len(mcp_rows))
|
|
self.assertEqual(len(cli_rows), len({item.name for item in cli_rows}))
|
|
self.assertEqual(
|
|
set((*ALL_TOOLS, *APPLICATION_TOOLS)),
|
|
{item.name for item in mcp_rows},
|
|
)
|
|
self.assertEqual(
|
|
{"read", "proposal", "application"},
|
|
{item.surface for item in mcp_rows},
|
|
)
|
|
|
|
def test_atomic_write_check_and_drift_detection(self) -> None:
|
|
content = generate_command_reference_bytes(ALPHA)
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
repository = Path(directory).resolve()
|
|
(repository / "generated").mkdir()
|
|
output = Path("generated/command-reference.md")
|
|
target = repository / output
|
|
|
|
self.assertEqual(
|
|
"written",
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=repository,
|
|
project_root=ALPHA,
|
|
output=output,
|
|
check=False,
|
|
),
|
|
)
|
|
self.assertEqual(content, target.read_bytes())
|
|
self.assertEqual(
|
|
"current",
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=repository,
|
|
project_root=ALPHA,
|
|
output=output,
|
|
check=True,
|
|
),
|
|
)
|
|
target.write_text("stale\n", encoding="utf-8")
|
|
with self.assertRaises(CommandReferenceDrift):
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=repository,
|
|
project_root=ALPHA,
|
|
output=output,
|
|
check=True,
|
|
)
|
|
self.assertEqual(b"stale\n", target.read_bytes())
|
|
self.assertEqual(
|
|
"written",
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=repository,
|
|
project_root=ALPHA,
|
|
output=output,
|
|
check=False,
|
|
),
|
|
)
|
|
self.assertEqual(content, target.read_bytes())
|
|
self.assertEqual(
|
|
"unchanged",
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=repository,
|
|
project_root=ALPHA,
|
|
output=output,
|
|
check=False,
|
|
),
|
|
)
|
|
|
|
def test_output_confinement_rejects_escapes_and_symlinks(self) -> None:
|
|
content = b"reference\n"
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
repository = Path(directory).resolve()
|
|
generated = repository / "generated"
|
|
generated.mkdir()
|
|
outside = repository.parent / f"{repository.name}-outside.md"
|
|
symlink = generated / "reference.md"
|
|
symlink.symlink_to(outside)
|
|
|
|
for output in (Path("../escape.md"), Path("generated/reference.md")):
|
|
with (
|
|
self.subTest(output=output),
|
|
self.assertRaises(CommandReferenceToolError),
|
|
):
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=repository,
|
|
project_root=ALPHA,
|
|
output=output,
|
|
check=False,
|
|
)
|
|
symlink.unlink()
|
|
(repository / "linked").symlink_to(generated, target_is_directory=True)
|
|
with self.assertRaises(CommandReferenceToolError):
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=repository,
|
|
project_root=ALPHA,
|
|
output=Path("linked/reference.md"),
|
|
check=False,
|
|
)
|
|
|
|
with self.assertRaises(CommandReferenceToolError):
|
|
publish_or_check_command_reference(
|
|
content,
|
|
repository_root=ROOT,
|
|
project_root=ALPHA,
|
|
output=Path("tests/fixtures/alpha/generated-reference.md"),
|
|
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()
|
|
(repository / "generated").mkdir()
|
|
output = Path("generated/reference.md")
|
|
stdout = io.StringIO()
|
|
|
|
with contextlib.redirect_stdout(stdout):
|
|
status = main(
|
|
(
|
|
"--repository-root",
|
|
str(repository),
|
|
"--project-root",
|
|
str(ALPHA),
|
|
"--output",
|
|
str(output),
|
|
)
|
|
)
|
|
|
|
self.assertEqual(0, status)
|
|
self.assertTrue((repository / output).is_file())
|
|
self.assertIn('"cli_rows":28', stdout.getvalue())
|
|
self.assertIn('"mcp_rows":36', stdout.getvalue())
|
|
|
|
target = repository / output
|
|
target.write_text("drift\n", encoding="utf-8")
|
|
stderr = io.StringIO()
|
|
with contextlib.redirect_stderr(stderr):
|
|
check_status = main(
|
|
(
|
|
"--repository-root",
|
|
str(repository),
|
|
"--project-root",
|
|
str(ALPHA),
|
|
"--output",
|
|
str(output),
|
|
"--check",
|
|
)
|
|
)
|
|
|
|
self.assertEqual(1, check_status)
|
|
self.assertEqual("drift\n", target.read_text(encoding="utf-8"))
|
|
self.assertIn("missing or stale", stderr.getvalue())
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|