diff --git a/tests/test_command_reference_tool.py b/tests/test_command_reference_tool.py new file mode 100644 index 0000000..a6c4feb --- /dev/null +++ b/tests/test_command_reference_tool.py @@ -0,0 +1,232 @@ +from __future__ import annotations + +import asyncio +import contextlib +import hashlib +import io +import os +import tempfile +import unittest +from pathlib import Path + +from docforge.mcp_server import ALL_TOOLS, APPLICATION_TOOLS +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_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() diff --git a/tools/generate_command_reference.py b/tools/generate_command_reference.py new file mode 100644 index 0000000..e448467 --- /dev/null +++ b/tools/generate_command_reference.py @@ -0,0 +1,369 @@ +"""Generate or verify the deterministic DocForge CLI and MCP command reference.""" + +from __future__ import annotations + +import argparse +import asyncio +import contextlib +import json +import os +import secrets +import shutil +import stat +import sys +import tempfile +from collections.abc import Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Literal + +from mcp.shared.memory import create_connected_server_and_client_session + +from docforge.command_reference import ( + CliCommandReference, + McpToolReference, + cli_command_references, + mcp_tool_references, + render_command_reference_markdown, +) +from docforge.index import ProjectIndex +from docforge.mcp_server import ALL_TOOLS, APPLICATION_TOOLS, create_server +from docforge.project import Project + +REPOSITORY_ROOT = Path(__file__).resolve().parents[1] +DEFAULT_PROJECT_ROOT = REPOSITORY_ROOT / "tests" / "fixtures" / "alpha" +EXPECTED_CLI_ROWS = 28 +EXPECTED_MCP_ROWS = 36 +MAX_REFERENCE_BYTES = 5_000_000 + + +class CommandReferenceToolError(RuntimeError): + """One safe repository-tool validation or publication failure.""" + + +class CommandReferenceDrift(CommandReferenceToolError): + """The checked output does not match current registered command metadata.""" + + +@dataclass(frozen=True) +class _FileIdentity: + device: int + inode: int + mode: int + size: int + modified_ns: int + changed_ns: int + + +async def collect_command_reference_rows( + project_root: Path, + *, + proposal_writer: str, + canonical_applier: str, +) -> tuple[tuple[CliCommandReference, ...], tuple[McpToolReference, ...]]: + """Collect the real CLI parser and full project-bound MCP registration surface.""" + + source = _safe_project_root(project_root) + with tempfile.TemporaryDirectory(prefix="docforge-command-reference-") as directory: + copied = Path(directory) / "project" + shutil.copytree(source, copied, symlinks=True) + ProjectIndex(Project.open(copied)).build() + server = create_server( + copied, + proposal_writer, + canonical_applier_id=canonical_applier, + ) + async with create_connected_server_and_client_session( + server, + raise_exceptions=True, + ) as session: + tools = tuple((await session.list_tools()).tools) + + cli_rows = cli_command_references() + expected_mcp_names = (*ALL_TOOLS, *APPLICATION_TOOLS) + mcp_rows = mcp_tool_references(tools, expected_names=expected_mcp_names) + if len(cli_rows) != EXPECTED_CLI_ROWS: + raise CommandReferenceToolError( + f"Expected {EXPECTED_CLI_ROWS} CLI rows, found {len(cli_rows)}" + ) + if len(mcp_rows) != EXPECTED_MCP_ROWS: + raise CommandReferenceToolError( + f"Expected {EXPECTED_MCP_ROWS} MCP rows, found {len(mcp_rows)}" + ) + return cli_rows, mcp_rows + + +def generate_command_reference_bytes( + project_root: Path, + *, + proposal_writer: str = "alpha-editor", + canonical_applier: str = "alpha-editor", +) -> bytes: + """Render one stable Markdown document from an isolated project copy.""" + + cli_rows, mcp_rows = asyncio.run( + collect_command_reference_rows( + project_root, + proposal_writer=proposal_writer, + canonical_applier=canonical_applier, + ) + ) + rendered = render_command_reference_markdown(cli_rows, mcp_rows).encode("utf-8") + if len(rendered) > MAX_REFERENCE_BYTES: + raise CommandReferenceToolError("Generated command reference exceeds its byte limit") + return rendered + + +def publish_or_check_command_reference( + content: bytes, + *, + repository_root: Path, + project_root: Path, + output: Path, + check: bool, +) -> Literal["current", "unchanged", "written"]: + """Safely check or atomically publish one repository-confined Markdown file.""" + + if len(content) > MAX_REFERENCE_BYTES: + raise CommandReferenceToolError("Command reference exceeds its byte limit") + root = _safe_repository_root(repository_root) + source = _safe_project_root(project_root) + relative = _safe_output_relative(root, source, output) + parent_parts = relative.parts[:-1] + target_name = relative.name + parent_fd = _open_relative_directory(root, parent_parts) + try: + initial = _file_identity(parent_fd, target_name) + existing = ( + _read_regular_file(parent_fd, target_name, limit=MAX_REFERENCE_BYTES) + if initial is not None + else None + ) + if _file_identity(parent_fd, target_name) != initial: + raise CommandReferenceToolError("Output target changed during inspection") + if check: + if existing != content: + raise CommandReferenceDrift( + f"Command reference is missing or stale: {relative.as_posix()}" + ) + return "current" + if existing == content: + return "unchanged" + _atomic_replace(parent_fd, target_name, content, expected=initial) + return "written" + finally: + os.close(parent_fd) + + +def _safe_repository_root(path: Path) -> Path: + if path.is_symlink(): + raise CommandReferenceToolError("Repository root must not be a symbolic link") + try: + resolved = path.resolve(strict=True) + except OSError as error: + raise CommandReferenceToolError("Repository root does not exist") from error + if not resolved.is_dir(): + raise CommandReferenceToolError("Repository root must be a directory") + return resolved + + +def _safe_project_root(path: Path) -> Path: + if path.is_symlink(): + raise CommandReferenceToolError("Project fixture root must not be a symbolic link") + try: + resolved = path.resolve(strict=True) + except OSError as error: + raise CommandReferenceToolError("Project fixture root does not exist") from error + if not resolved.is_dir(): + raise CommandReferenceToolError("Project fixture root must be a directory") + return resolved + + +def _safe_output_relative(root: Path, project_root: Path, output: Path) -> Path: + candidate = output if output.is_absolute() else root / output + absolute = Path(os.path.abspath(candidate)) + if absolute == root or not absolute.is_relative_to(root): + raise CommandReferenceToolError("Output must remain inside the repository root") + if absolute == project_root or absolute.is_relative_to(project_root): + raise CommandReferenceToolError("Output must not modify the source project fixture") + relative = absolute.relative_to(root) + if relative.suffix != ".md": + raise CommandReferenceToolError("Output must be one Markdown file") + if not relative.name or len(relative.parts) < 2: + raise CommandReferenceToolError("Output must be below an existing repository directory") + return relative + + +def _open_relative_directory(root: Path, parts: tuple[str, ...]) -> int: + flags = os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW + descriptor = os.open(root, flags) + try: + for part in parts: + if part in {"", ".", ".."}: + raise CommandReferenceToolError("Output directory is unsafe") + try: + next_descriptor = os.open(part, flags, dir_fd=descriptor) + except OSError as error: + raise CommandReferenceToolError( + "Output parent must be an existing non-symlink directory" + ) from error + os.close(descriptor) + descriptor = next_descriptor + return descriptor + except Exception: + os.close(descriptor) + raise + + +def _file_identity(directory_fd: int, name: str) -> _FileIdentity | None: + try: + status = os.stat(name, dir_fd=directory_fd, follow_symlinks=False) + except FileNotFoundError: + return None + if not stat.S_ISREG(status.st_mode): + raise CommandReferenceToolError("Output target must be a regular file") + return _FileIdentity( + device=status.st_dev, + inode=status.st_ino, + mode=status.st_mode, + size=status.st_size, + modified_ns=status.st_mtime_ns, + changed_ns=status.st_ctime_ns, + ) + + +def _read_regular_file(directory_fd: int, name: str, *, limit: int) -> bytes: + try: + descriptor = os.open(name, os.O_RDONLY | os.O_NOFOLLOW, dir_fd=directory_fd) + except OSError as error: + raise CommandReferenceToolError("Output target could not be opened safely") from error + try: + before = os.fstat(descriptor) + if not stat.S_ISREG(before.st_mode) or before.st_size > limit: + raise CommandReferenceToolError("Output target is not a bounded regular file") + chunks: list[bytes] = [] + remaining = limit + 1 + while remaining: + chunk = os.read(descriptor, min(remaining, 64 * 1024)) + if not chunk: + break + chunks.append(chunk) + remaining -= len(chunk) + content = b"".join(chunks) + after = os.fstat(descriptor) + if len(content) > limit or _identity_from_stat(before) != _identity_from_stat(after): + raise CommandReferenceToolError("Output target changed during inspection") + return content + finally: + os.close(descriptor) + + +def _identity_from_stat(status: os.stat_result) -> _FileIdentity: + return _FileIdentity( + device=status.st_dev, + inode=status.st_ino, + mode=status.st_mode, + size=status.st_size, + modified_ns=status.st_mtime_ns, + changed_ns=status.st_ctime_ns, + ) + + +def _atomic_replace( + directory_fd: int, + name: str, + content: bytes, + *, + expected: _FileIdentity | None, +) -> None: + temporary_name = f".{name}.docforge-command-reference-{secrets.token_hex(12)}" + flags = os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW + try: + descriptor = os.open(temporary_name, flags, 0o644, dir_fd=directory_fd) + except OSError as error: + raise CommandReferenceToolError("Could not create atomic output temporary") from error + try: + view = memoryview(content) + while view: + written = os.write(descriptor, view) + if written <= 0: + raise CommandReferenceToolError("Could not write command reference output") + view = view[written:] + os.fsync(descriptor) + except Exception: + os.close(descriptor) + _unlink_at(directory_fd, temporary_name) + raise + else: + os.close(descriptor) + try: + if _file_identity(directory_fd, name) != expected: + raise CommandReferenceToolError("Output target changed before atomic publication") + os.replace( + temporary_name, + name, + src_dir_fd=directory_fd, + dst_dir_fd=directory_fd, + ) + os.fsync(directory_fd) + except Exception: + _unlink_at(directory_fd, temporary_name) + raise + + +def _unlink_at(directory_fd: int, name: str) -> None: + with contextlib.suppress(FileNotFoundError): + os.unlink(name, dir_fd=directory_fd) + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="Generate or check deterministic DocForge command-reference Markdown." + ) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--check", action="store_true") + parser.add_argument("--repository-root", type=Path, default=REPOSITORY_ROOT) + parser.add_argument("--project-root", type=Path, default=DEFAULT_PROJECT_ROOT) + parser.add_argument("--proposal-writer", default="alpha-editor") + parser.add_argument("--canonical-applier", default="alpha-editor") + return parser + + +def main(arguments: Sequence[str] | None = None) -> int: + parsed = _parser().parse_args(arguments) + try: + content = generate_command_reference_bytes( + parsed.project_root, + proposal_writer=parsed.proposal_writer, + canonical_applier=parsed.canonical_applier, + ) + state = publish_or_check_command_reference( + content, + repository_root=parsed.repository_root, + project_root=parsed.project_root, + output=parsed.output, + check=parsed.check, + ) + except CommandReferenceDrift as error: + print(str(error), file=sys.stderr) + return 1 + except Exception as error: + print(f"command-reference generation failed: {error}", file=sys.stderr) + return 2 + print( + json.dumps( + { + "status": state, + "output": str(parsed.output), + "cli_rows": EXPECTED_CLI_ROWS, + "mcp_rows": EXPECTED_MCP_ROWS, + }, + sort_keys=True, + separators=(",", ":"), + ) + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main())