1
0
Fork 0
Code Issues Pull requests Projects Releases 2 Packages Wiki Activity Actions Pages
DocForge2/tests/test_command_reference_tool.py

405 lines
15 KiB
Python
Raw Permalink Normal View History

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"))
2026-07-29 15:10:47 -04:00
self.assertTrue(first.startswith(b"# DocForge command reference\n"))
self.assertIn(b"Do not edit this file by hand.", first)
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()