Harden M7.3 server: typed schemas, depth tracking, fnmatch, config preservation, and collision detection

This commit is contained in:
rob 2026-07-20 01:14:04 -03:00
parent 3525e2eeae
commit fe53721725
6 changed files with 502 additions and 69 deletions

View File

@ -457,7 +457,10 @@ def main():
# mcp serve # mcp serve
p_mcp_serve = mcp_sub.add_parser("serve", help="Start CmdForge as an MCP server") p_mcp_serve = mcp_sub.add_parser("serve", help="Start CmdForge as an MCP server")
p_mcp_serve.add_argument("--transport", default="stdio", help="Transport (default: stdio)") p_mcp_serve.add_argument(
"--transport", choices=["stdio"], default="stdio",
help="Transport (M7.3 supports stdio)",
)
p_mcp_serve.set_defaults(func=cmd_mcp) p_mcp_serve.set_defaults(func=cmd_mcp)
# mcp list # mcp list

View File

@ -1,6 +1,8 @@
"""MCP client support built on the official MCP Python SDK.""" """MCP client support built on the official MCP Python SDK."""
import asyncio import asyncio
import contextvars
import fnmatch
import hashlib import hashlib
import json import json
import os import os
@ -16,6 +18,7 @@ import yaml
MCP_CONFIG_FILE = Path.home() / ".cmdforge" / "mcp.yaml" MCP_CONFIG_FILE = Path.home() / ".cmdforge" / "mcp.yaml"
MCP_CONFIG_VERSION = 1 MCP_CONFIG_VERSION = 1
MCP_DEPTH_ENV = "CMDFORGE_MCP_DEPTH"
RESULT_MODES = ("auto", "structured", "content", "text") RESULT_MODES = ("auto", "structured", "content", "text")
SUPPORTED_TRANSPORTS = ("stdio",) SUPPORTED_TRANSPORTS = ("stdio",)
DEFAULT_INHERITED_ENV = ( DEFAULT_INHERITED_ENV = (
@ -41,6 +44,9 @@ _SECRET_PATTERNS = (
) )
T = TypeVar("T") T = TypeVar("T")
_MCP_CALL_DEPTH: contextvars.ContextVar[int] = contextvars.ContextVar(
"cmdforge_mcp_call_depth", default=0
)
@dataclass @dataclass
@ -50,6 +56,15 @@ class McpServeConfig:
expose: List[str] = field(default_factory=list) expose: List[str] = field(default_factory=list)
deny: List[str] = field(default_factory=list) deny: List[str] = field(default_factory=list)
def __post_init__(self) -> None:
for field_name, patterns in (("expose", self.expose), ("deny", self.deny)):
if not isinstance(patterns, list) or not all(
isinstance(pattern, str) and pattern for pattern in patterns
):
raise ValueError(
f"MCP server '{field_name}' must be a list of non-empty strings"
)
def is_exposed(self, tool_name: str) -> bool: def is_exposed(self, tool_name: str) -> bool:
"""Check whether a tool should be exposed to MCP clients.""" """Check whether a tool should be exposed to MCP clients."""
if self._matches_deny(tool_name): if self._matches_deny(tool_name):
@ -72,14 +87,8 @@ class McpServeConfig:
def _glob_match(pattern: str, name: str) -> bool: def _glob_match(pattern: str, name: str) -> bool:
"""Simple wildcard matching: * matches any sequence, name is literal.""" """Match a policy pattern using shell-style wildcards."""
if pattern == "*": return fnmatch.fnmatchcase(name, pattern)
return True
if pattern == name:
return True
if pattern.endswith("-*") and name.startswith(pattern[:-1]):
return True
return False
@dataclass @dataclass
@ -188,19 +197,26 @@ def _build_server_env(cfg: McpServerConfig) -> Dict[str, str]:
} }
for name, value in cfg.env.items(): for name, value in cfg.env.items():
environment[name] = _expand_env_value(value) environment[name] = _expand_env_value(value)
call_depth = _MCP_CALL_DEPTH.get()
if call_depth > 0:
environment[MCP_DEPTH_ENV] = str(call_depth)
return environment return environment
def load_mcp_config() -> List[McpServerConfig]: def _load_mcp_document() -> dict:
if not MCP_CONFIG_FILE.exists(): if not MCP_CONFIG_FILE.exists():
return [] return {}
data = yaml.safe_load(MCP_CONFIG_FILE.read_text()) or {} data = yaml.safe_load(MCP_CONFIG_FILE.read_text()) or {}
if not isinstance(data, dict): if not isinstance(data, dict):
raise ValueError("MCP config must be a YAML mapping") raise ValueError("MCP config must be a YAML mapping")
version = data.get("version", MCP_CONFIG_VERSION) version = data.get("version", MCP_CONFIG_VERSION)
if version != MCP_CONFIG_VERSION: if version != MCP_CONFIG_VERSION:
raise ValueError(f"Unsupported MCP config version: {version}") raise ValueError(f"Unsupported MCP config version: {version}")
return data
def load_mcp_config() -> List[McpServerConfig]:
data = _load_mcp_document()
raw_servers = data.get("servers") or {} raw_servers = data.get("servers") or {}
if not isinstance(raw_servers, dict): if not isinstance(raw_servers, dict):
raise ValueError("MCP config 'servers' must be a mapping") raise ValueError("MCP config 'servers' must be a mapping")
@ -230,10 +246,9 @@ def save_mcp_config(servers: List[McpServerConfig]) -> None:
for server in servers: for server in servers:
server.validate() server.validate()
MCP_CONFIG_FILE.parent.mkdir(parents=True, exist_ok=True, mode=0o700) MCP_CONFIG_FILE.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
body = { body = _load_mcp_document()
"version": MCP_CONFIG_VERSION, body["version"] = MCP_CONFIG_VERSION
"servers": {server.name: _server_to_dict(server) for server in servers}, body["servers"] = {server.name: _server_to_dict(server) for server in servers}
}
serialized = yaml.safe_dump(body, default_flow_style=False, sort_keys=False) serialized = yaml.safe_dump(body, default_flow_style=False, sort_keys=False)
temp_path = None temp_path = None
try: try:
@ -289,14 +304,10 @@ def _require_mcp_sdk() -> None:
def load_mcp_serve_config() -> McpServeConfig: def load_mcp_serve_config() -> McpServeConfig:
if not MCP_CONFIG_FILE.exists(): data = _load_mcp_document()
return McpServeConfig()
data = yaml.safe_load(MCP_CONFIG_FILE.read_text()) or {}
if not isinstance(data, dict):
return McpServeConfig()
server_raw = data.get("server") or {} server_raw = data.get("server") or {}
if not isinstance(server_raw, dict): if not isinstance(server_raw, dict):
return McpServeConfig() raise ValueError("MCP config 'server' must be a mapping")
return McpServeConfig( return McpServeConfig(
expose=server_raw.get("expose", []), expose=server_raw.get("expose", []),
deny=server_raw.get("deny", []), deny=server_raw.get("deny", []),

View File

@ -1,15 +1,21 @@
"""MCP server support: expose CmdForge tools as MCP tools.""" """MCP server support: expose CmdForge tools as MCP tools."""
import inspect
import os
import sys import sys
from typing import Any from typing import Annotated, Any, Literal
from .mcp_client import McpServeConfig, load_mcp_serve_config, _require_mcp_sdk from .mcp_client import (
MCP_DEPTH_ENV,
McpServeConfig,
_MCP_CALL_DEPTH,
_require_mcp_sdk,
load_mcp_serve_config,
)
from .runner import MAX_TOOL_DEPTH
from .tool import ToolArgument from .tool import ToolArgument
MAX_TOOL_DEPTH = 10
def _cmdforge_tool_name(name: str) -> str: def _cmdforge_tool_name(name: str) -> str:
"""Map a CmdForge qualified tool name to an MCP-safe name. """Map a CmdForge qualified tool name to an MCP-safe name.
@ -21,7 +27,7 @@ def _cmdforge_tool_name(name: str) -> str:
def _argument_to_json_schema(arg: ToolArgument) -> dict: def _argument_to_json_schema(arg: ToolArgument) -> dict:
"""Map a CmdForge ToolArgument to a JSON Schema property.""" """Map a CmdForge ToolArgument to a JSON Schema property."""
prop: dict = {} prop: dict = {}
arg_type = getattr(arg, "type", None) or "string" arg_type = arg.type
type_map = { type_map = {
"string": "string", "string": "string",
@ -31,15 +37,16 @@ def _argument_to_json_schema(arg: ToolArgument) -> dict:
} }
prop["type"] = type_map.get(arg_type, "string") prop["type"] = type_map.get(arg_type, "string")
if getattr(arg, "description", ""): if arg.description:
prop["description"] = arg.description prop["description"] = arg.description
if getattr(arg, "default", None) is not None: if arg.default is not None:
prop["default"] = arg.default prop["default"] = _coerce_argument_value(arg.default, arg.type)
enum_vals = getattr(arg, "enum", None) if arg.enum:
if enum_vals: prop["enum"] = [
prop["enum"] = enum_vals _coerce_argument_value(value, arg.type) for value in arg.enum
]
return prop return prop
@ -56,9 +63,16 @@ def _build_tool_schema(tool) -> dict:
for arg in tool.arguments: for arg in tool.arguments:
prop = _argument_to_json_schema(arg) prop = _argument_to_json_schema(arg)
properties[arg.variable] = prop properties[arg.variable] = prop
if getattr(arg, "required", False): if arg.required:
required.append(arg.variable) required.append(arg.variable)
if "input" not in properties:
properties["input"] = {
"type": "string",
"description": "Text passed to the CmdForge tool as standard input.",
"default": "",
}
input_schema = { input_schema = {
"type": "object", "type": "object",
"properties": properties, "properties": properties,
@ -90,12 +104,11 @@ def serve(transport: str = "stdio") -> None:
serve_config = load_mcp_serve_config() serve_config = load_mcp_serve_config()
from .tool import list_tools, load_tool from .tool import list_tools, load_tool
from .runner import MAX_TOOL_DEPTH as _MAX_TOOL_DEPTH
server = FastMCP("CmdForge") server = FastMCP("CmdForge")
tool_names = list_tools() tool_names = list_tools()
exposed_count = 0 exposed_count = 0
exposed_names = {}
for name in sorted(tool_names): for name in sorted(tool_names):
# Skip private/unlisted tools unless explicitly exposed # Skip private/unlisted tools unless explicitly exposed
@ -103,10 +116,17 @@ def serve(transport: str = "stdio") -> None:
if not tool_obj: if not tool_obj:
continue continue
if not _is_exposable(tool_obj, serve_config): if not _is_exposable(tool_obj, serve_config, qualified_name=name):
continue continue
mcp_name = _cmdforge_tool_name(name) mcp_name = _cmdforge_tool_name(name)
previous = exposed_names.get(mcp_name)
if previous is not None:
raise RuntimeError(
f"MCP tool name collision: '{previous}' and '{name}' both map to "
f"'{mcp_name}'"
)
exposed_names[mcp_name] = name
schema = _build_tool_schema(tool_obj) schema = _build_tool_schema(tool_obj)
# Capture tool reference for the handler # Capture tool reference for the handler
@ -119,44 +139,169 @@ def serve(transport: str = "stdio") -> None:
server.run(transport=transport) server.run(transport=transport)
def _is_exposable(tool, config: McpServeConfig) -> bool: def _is_exposable(
tool, config: McpServeConfig, qualified_name: str | None = None
) -> bool:
"""Check whether a tool should be exposed per the server policy. """Check whether a tool should be exposed per the server policy.
Private/unlisted tools must be explicitly listed in expose. Public tools Private/unlisted tools must be explicitly listed in expose. Public tools
are exposed only if the expose list includes them or '*'. are exposed only if the expose list includes them or '*'.
""" """
policy_name = qualified_name or tool.name
if config._matches_deny(policy_name):
return False
if tool.visibility not in ("public", None, ""): if tool.visibility not in ("public", None, ""):
# private or unlisted — must be explicitly named # private or unlisted — must be explicitly named
return tool.name in config.expose return policy_name in config.expose
return config.is_exposed(tool.name) return config.is_exposed(policy_name)
def _register_tool(server, cmdforge_name: str, mcp_name: str, schema: dict, tool_obj) -> None: def _register_tool(server, cmdforge_name: str, mcp_name: str, schema: dict, tool_obj) -> None:
"""Register a CmdForge tool handler with the FastMCP server.""" """Register a CmdForge tool handler with the FastMCP server."""
tool_desc = schema.get("description", "") tool_desc = schema.get("description", "")
input_schema = schema.get("inputSchema", {}) argument_names = {arg.variable for arg in tool_obj.arguments}
@server.tool(
name=mcp_name,
description=tool_desc,
)
def handler(**kwargs: Any) -> str: def handler(**kwargs: Any) -> str:
from .runner import run_tool from .runner import run_tool
inherited_depth = _read_mcp_depth()
if inherited_depth >= MAX_TOOL_DEPTH:
raise RuntimeError(
f"Maximum MCP nesting depth ({MAX_TOOL_DEPTH}) exceeded"
)
if "input" in argument_names:
input_text = kwargs.get("input", "") or ""
else:
input_text = kwargs.pop("input", "") or ""
token = _MCP_CALL_DEPTH.set(inherited_depth + 1)
try:
output, exit_code = run_tool( output, exit_code = run_tool(
tool_obj, tool_obj,
input_text="", input_text=input_text,
custom_args=kwargs, custom_args=kwargs,
provider_override=None, provider_override=None,
verbose=False, verbose=False,
dry_run=False, dry_run=False,
_depth=MAX_TOOL_DEPTH, _depth=inherited_depth,
) )
finally:
_MCP_CALL_DEPTH.reset(token)
if exit_code != 0: if exit_code != 0:
raise RuntimeError(f"Tool '{cmdforge_name}' exited with code {exit_code}") raise RuntimeError(f"Tool '{cmdforge_name}' exited with code {exit_code}")
return output return output
# FastMCP extracts the schema from type hints. We also set the input_schema handler.__name__ = f"cmdforge_{mcp_name}"
# via an internal mechanism for richer descriptions. The `**kwargs` handler handler.__signature__ = _build_handler_signature(tool_obj)
# receives all tool arguments as named parameters. server.add_tool(
handler,
name=mcp_name,
description=tool_desc,
structured_output=False,
)
def _build_handler_signature(tool) -> inspect.Signature:
"""Create the flat typed signature FastMCP uses for schema and validation."""
required_parameters = []
optional_parameters = []
seen = set()
for arg in tool.arguments:
if not arg.variable.isidentifier():
raise ValueError(
f"Tool '{tool.name}' has invalid argument variable '{arg.variable}'"
)
if arg.variable in seen:
raise ValueError(
f"Tool '{tool.name}' defines argument '{arg.variable}' more than once"
)
seen.add(arg.variable)
annotation = _argument_annotation(arg)
default = inspect.Parameter.empty
if not arg.required:
default = (
_coerce_argument_value(arg.default, arg.type)
if arg.default is not None
else None
)
parameter = inspect.Parameter(
arg.variable,
kind=inspect.Parameter.KEYWORD_ONLY,
default=default,
annotation=annotation,
)
(required_parameters if arg.required else optional_parameters).append(parameter)
if "input" not in seen:
optional_parameters.append(
inspect.Parameter(
"input",
kind=inspect.Parameter.KEYWORD_ONLY,
default="",
annotation=Annotated[
str,
_pydantic_field(
"Text passed to the CmdForge tool as standard input."
),
],
)
)
return inspect.Signature(
required_parameters + optional_parameters,
return_annotation=str,
)
def _argument_annotation(arg: ToolArgument):
python_type = {
"string": str,
"integer": int,
"number": float,
"boolean": bool,
}[arg.type]
if arg.enum:
values = tuple(_coerce_argument_value(value, arg.type) for value in arg.enum)
python_type = Literal[values]
if arg.description:
return Annotated[python_type, _pydantic_field(arg.description)]
return python_type
def _pydantic_field(description: str):
from pydantic import Field
return Field(description=description)
def _coerce_argument_value(value: Any, arg_type: str) -> Any:
if value is None:
return None
if arg_type == "string":
return str(value)
if arg_type == "integer":
return int(value)
if arg_type == "number":
return float(value)
if arg_type == "boolean":
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in ("true", "1", "yes", "on"):
return True
if normalized in ("false", "0", "no", "off"):
return False
raise ValueError(f"Invalid boolean value: {value}")
return bool(value)
raise ValueError(f"Unsupported argument type: {arg_type}")
def _read_mcp_depth() -> int:
raw_depth = os.environ.get(MCP_DEPTH_ENV, "0")
try:
depth = int(raw_depth)
except ValueError as exc:
raise RuntimeError(f"Invalid {MCP_DEPTH_ENV} value: {raw_depth}") from exc
if depth < 0:
raise RuntimeError(f"Invalid {MCP_DEPTH_ENV} value: {raw_depth}")
return depth

View File

@ -5,7 +5,7 @@ import shutil
import stat import stat
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Optional, List, Literal from typing import Any, Optional, List, Literal
import yaml import yaml
@ -70,15 +70,36 @@ class ToolArgument:
"""Definition of a custom input argument.""" """Definition of a custom input argument."""
flag: str # e.g., "--max-size" flag: str # e.g., "--max-size"
variable: str # e.g., "max_size" variable: str # e.g., "max_size"
default: Optional[str] = None default: Optional[Any] = None
description: str = "" description: str = ""
type: str = "string"
enum: Optional[List[Any]] = None
required: bool = False
def __post_init__(self) -> None:
if self.type not in ("string", "integer", "number", "boolean"):
raise ValueError(
"ToolArgument type must be string, integer, number, or boolean"
)
if self.enum is not None and (
not isinstance(self.enum, list) or not self.enum
):
raise ValueError("ToolArgument enum must be a non-empty list")
if not isinstance(self.required, bool):
raise ValueError("ToolArgument required must be a boolean")
def to_dict(self) -> dict: def to_dict(self) -> dict:
d = {"flag": self.flag, "variable": self.variable} d = {"flag": self.flag, "variable": self.variable}
if self.default: if self.default is not None:
d["default"] = self.default d["default"] = self.default
if self.description: if self.description:
d["description"] = self.description d["description"] = self.description
if self.type != "string":
d["type"] = self.type
if self.enum is not None:
d["enum"] = self.enum
if self.required:
d["required"] = True
return d return d
@classmethod @classmethod
@ -87,7 +108,10 @@ class ToolArgument:
flag=data["flag"], flag=data["flag"],
variable=data["variable"], variable=data["variable"],
default=data.get("default"), default=data.get("default"),
description=data.get("description", "") description=data.get("description", ""),
type=data.get("type", "string"),
enum=data.get("enum"),
required=data.get("required", False),
) )

View File

@ -5,7 +5,8 @@ import time
from pathlib import Path from pathlib import Path
import pytest import pytest
from unittest.mock import patch, MagicMock import yaml
from unittest.mock import patch
from cmdforge.mcp_client import ( from cmdforge.mcp_client import (
DEFAULT_INHERITED_ENV, DEFAULT_INHERITED_ENV,
McpServerConfig, McpServerConfig,
@ -126,6 +127,17 @@ class TestMcpConfigPersistence:
perms = oct(temp_mcp_file.stat().st_mode & 0o777) perms = oct(temp_mcp_file.stat().st_mode & 0o777)
assert perms == "0o600" assert perms == "0o600"
def test_save_preserves_server_exposure_policy(self, temp_mcp_file):
temp_mcp_file.write_text(
"version: 1\nserver:\n expose: [summarize]\n deny: [dangerous]\n"
)
save_mcp_config([McpServerConfig(name="test", command="echo")])
data = yaml.safe_load(temp_mcp_file.read_text())
assert data["server"] == {
"expose": ["summarize"],
"deny": ["dangerous"],
}
class TestResultNormalization: class TestResultNormalization:
def test_auto_prefers_structured_content(self): def test_auto_prefers_structured_content(self):
@ -432,6 +444,50 @@ class TestMcpServeConfig:
assert config.is_exposed("extract") assert config.is_exposed("extract")
assert not config.is_exposed("other") assert not config.is_exposed("other")
@pytest.mark.parametrize(
"field, value",
[("expose", "*"), ("deny", "dangerous"), ("expose", [""])],
)
def test_rejects_malformed_policy_lists(self, field, value):
from cmdforge.mcp_client import McpServeConfig
values = {"expose": [], "deny": [], field: value}
with pytest.raises(ValueError, match="list of non-empty strings"):
McpServeConfig(**values)
def test_general_wildcard_matching(self):
from cmdforge.mcp_client import McpServeConfig
config = McpServeConfig(expose=["official/*"], deny=["*/dangerous"])
assert config.is_exposed("official/summarize")
assert not config.is_exposed("official/dangerous")
def test_private_tool_still_honors_deny(self):
from cmdforge.mcp_client import McpServeConfig
from cmdforge.mcp_server import _is_exposable
from cmdforge.tool import Tool
tool = Tool(name="dangerous", visibility="private")
config = McpServeConfig(expose=["official/dangerous"], deny=["official/*"])
assert not _is_exposable(tool, config, qualified_name="official/dangerous")
def test_private_qualified_tool_requires_qualified_expose(self):
from cmdforge.mcp_client import McpServeConfig
from cmdforge.mcp_server import _is_exposable
from cmdforge.tool import Tool
tool = Tool(name="extract", visibility="private")
assert not _is_exposable(
tool,
McpServeConfig(expose=["extract"]),
qualified_name="official/extract",
)
assert _is_exposable(
tool,
McpServeConfig(expose=["official/extract"]),
qualified_name="official/extract",
)
class TestToolNameMapping: class TestToolNameMapping:
def test_simple_name_passes_through(self): def test_simple_name_passes_through(self):
@ -455,24 +511,26 @@ class TestArgumentToJsonSchema:
def test_integer_arg_with_default(self): def test_integer_arg_with_default(self):
from cmdforge.tool import ToolArgument from cmdforge.tool import ToolArgument
from cmdforge.mcp_server import _argument_to_json_schema from cmdforge.mcp_server import _argument_to_json_schema
arg = ToolArgument(flag="--limit", variable="limit", default="10") arg = ToolArgument(
setattr(arg, "type", "integer") flag="--limit", variable="limit", default="10", type="integer"
)
schema = _argument_to_json_schema(arg) schema = _argument_to_json_schema(arg)
assert schema["type"] == "integer" assert schema["type"] == "integer"
assert schema["default"] == "10" assert schema["default"] == 10
def test_arg_with_enum(self): def test_arg_with_enum(self):
from cmdforge.tool import ToolArgument from cmdforge.tool import ToolArgument
from cmdforge.mcp_server import _argument_to_json_schema from cmdforge.mcp_server import _argument_to_json_schema
arg = ToolArgument(flag="--mode", variable="mode") arg = ToolArgument(
setattr(arg, "enum", ["fast", "accurate"]) flag="--mode", variable="mode", enum=["fast", "accurate"]
)
schema = _argument_to_json_schema(arg) schema = _argument_to_json_schema(arg)
assert schema["enum"] == ["fast", "accurate"] assert schema["enum"] == ["fast", "accurate"]
class TestBuildToolSchema: class TestBuildToolSchema:
def test_builds_input_schema(self, tmp_path): def test_builds_input_schema(self, tmp_path):
from cmdforge.tool import Tool, ToolArgument, save_tool from cmdforge.tool import Tool, ToolArgument
from cmdforge.mcp_server import _build_tool_schema from cmdforge.mcp_server import _build_tool_schema
with patch('cmdforge.tool.TOOLS_DIR', tmp_path / ".cmdforge"): with patch('cmdforge.tool.TOOLS_DIR', tmp_path / ".cmdforge"):
@ -494,3 +552,170 @@ class TestBuildToolSchema:
assert schema["inputSchema"]["type"] == "object" assert schema["inputSchema"]["type"] == "object"
assert "name" in schema["inputSchema"]["properties"] assert "name" in schema["inputSchema"]["properties"]
assert "count" in schema["inputSchema"]["properties"] assert "count" in schema["inputSchema"]["properties"]
assert "input" in schema["inputSchema"]["properties"]
class TestRegisteredMcpTool:
def test_fastmcp_uses_flat_typed_schema_and_invokes_tool(self):
import asyncio
from mcp.server.fastmcp import FastMCP
from cmdforge.mcp_server import _build_tool_schema, _register_tool
from cmdforge.tool import Tool, ToolArgument
server = FastMCP("test")
tool = Tool(
name="greet",
arguments=[
ToolArgument(
flag="--name",
variable="name",
type="string",
required=True,
description="Who to greet",
),
ToolArgument(
flag="--count",
variable="count",
type="integer",
default=1,
enum=[1, 2, 3],
),
],
output="Hello {name} x{count}: {input}",
)
_register_tool(server, "greet", "greet", _build_tool_schema(tool), tool)
async def exercise():
listed = (await server.list_tools())[0]
result = await server.call_tool(
"greet", {"name": "Alice", "count": 2, "input": "welcome"}
)
return listed, result
listed, result = asyncio.run(exercise())
assert listed.inputSchema["required"] == ["name"]
assert listed.inputSchema["properties"]["count"]["type"] == "integer"
assert listed.inputSchema["properties"]["count"]["enum"] == [1, 2, 3]
assert "kwargs" not in listed.inputSchema["properties"]
assert result[0].text == "Hello Alice x2: welcome"
def test_handler_starts_at_inherited_depth(self):
import asyncio
from mcp.server.fastmcp import FastMCP
from cmdforge.mcp_server import _build_tool_schema, _register_tool
from cmdforge.tool import Tool
server = FastMCP("test")
tool = Tool(name="depth", output="ok")
with patch("cmdforge.runner.run_tool", return_value=("ok", 0)) as run:
_register_tool(server, "depth", "depth", _build_tool_schema(tool), tool)
asyncio.run(server.call_tool("depth", {}))
assert run.call_args.kwargs["_depth"] == 0
def test_handler_rejects_exhausted_mcp_depth(self, monkeypatch):
import asyncio
from mcp.server.fastmcp import FastMCP
from cmdforge.mcp_server import (
MAX_TOOL_DEPTH,
_build_tool_schema,
_register_tool,
)
from cmdforge.tool import Tool
monkeypatch.setenv("CMDFORGE_MCP_DEPTH", str(MAX_TOOL_DEPTH))
server = FastMCP("test")
tool = Tool(name="depth", output="ok")
_register_tool(server, "depth", "depth", _build_tool_schema(tool), tool)
with pytest.raises(Exception, match="Maximum MCP nesting depth"):
asyncio.run(server.call_tool("depth", {}))
def test_serve_fails_closed_on_mapped_name_collision(self):
from cmdforge.mcp_client import McpServeConfig
from cmdforge.mcp_server import serve
from cmdforge.tool import Tool
tools = {
"owner/tool": Tool(name="tool"),
"owner__tool": Tool(name="owner__tool"),
}
with (
patch(
"cmdforge.mcp_server.load_mcp_serve_config",
return_value=McpServeConfig(expose=["*"]),
),
patch("cmdforge.tool.list_tools", return_value=list(tools)),
patch("cmdforge.tool.load_tool", side_effect=tools.get),
patch("mcp.server.fastmcp.FastMCP.run"),
):
with pytest.raises(RuntimeError, match="name collision"):
serve()
class TestCmdForgeMcpServerEndToEnd:
def test_stdio_discovery_and_invocation(self, tmp_path, monkeypatch):
pytest.importorskip("mcp")
home = tmp_path / "home"
cmdforge_dir = home / ".cmdforge"
cmdforge_dir.mkdir(parents=True)
(cmdforge_dir / "mcp.yaml").write_text(
yaml.safe_dump({
"version": 1,
"server": {
"expose": ["greet", "echo-input", "hidden"],
"deny": ["hidden"],
},
})
)
tools = {
"greet": {
"name": "greet",
"arguments": [{
"flag": "--name",
"variable": "name",
"type": "string",
"required": True,
}],
"steps": [],
"output": "Hello {name}",
},
"echo-input": {
"name": "echo-input",
"steps": [],
"output": "{input}",
},
"hidden": {
"name": "hidden",
"visibility": "private",
"steps": [],
"output": "secret",
},
}
for name, data in tools.items():
tool_dir = cmdforge_dir / name
tool_dir.mkdir()
(tool_dir / "config.yaml").write_text(yaml.safe_dump(data))
client_config = tmp_path / "client-mcp.yaml"
monkeypatch.setattr("cmdforge.mcp_client.MCP_CONFIG_FILE", client_config)
save_mcp_config([
McpServerConfig(
name="cmdforge",
command=sys.executable,
args=["-m", "cmdforge.cli", "mcp", "serve"],
env={"HOME": str(home)},
timeout=10,
approved=True,
)
])
manager = McpClientManager()
discovered = {tool["name"]: tool for tool in manager.discover("cmdforge")}
assert set(discovered) == {"greet", "echo-input"}
assert discovered["greet"]["inputSchema"]["required"] == ["name"]
assert manager.call_tool(
"cmdforge", "greet", {"name": "Alice"}, "text"
) == "Hello Alice"
assert manager.call_tool(
"cmdforge", "echo-input", {"input": "through stdin"}, "text"
) == "through stdin"

View File

@ -80,6 +80,31 @@ class TestToolArgument:
assert restored.default == original.default assert restored.default == original.default
assert restored.description == original.description assert restored.description == original.description
def test_typed_metadata_roundtrip(self):
original = ToolArgument(
flag="--limit",
variable="limit",
default=10,
description="Maximum results",
type="integer",
enum=[10, 25, 50],
required=True,
)
assert ToolArgument.from_dict(original.to_dict()) == original
@pytest.mark.parametrize(
"values",
[
{"type": "object"},
{"enum": []},
{"enum": "fast"},
{"required": "yes"},
],
)
def test_rejects_invalid_typed_metadata(self, values):
with pytest.raises(ValueError):
ToolArgument(flag="--value", variable="value", **values)
class TestPromptStep: class TestPromptStep:
"""Tests for PromptStep dataclass.""" """Tests for PromptStep dataclass."""