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.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()