mirror of
https://github.com/ggozad/oterm.git
synced 2026-10-10 09:13:20 +02:00
Tool identity was the bare tool name, so a name exported by two MCP
servers resolved to both toolsets and pydantic-ai rejected the agent.
MCP tools are now identified as `{server}_{tool}` and reach the model
through PrefixedToolset, making per-server selection independent.
Since the server name now reaches the provider, reject names outside
[a-zA-Z0-9_-] when the config loads and skip tools whose qualified name
is over 64 characters.
Tools load before the store so the 0.22.0 upgrade can qualify saved
selections from the connected servers' tool lists.
Point the ty pre-commit hook at the pinned version, matching CI and the
ruff hooks.
Closes #321
273 lines
9.9 KiB
Python
273 lines
9.9 KiB
Python
import pytest
|
|
from fastmcp.client.transports import (
|
|
SSETransport,
|
|
StdioTransport,
|
|
StreamableHttpTransport,
|
|
)
|
|
from pydantic_ai.mcp import MCPToolset
|
|
|
|
from oterm.tools.mcp.setup import (
|
|
_build_toolsets,
|
|
mcp_servers,
|
|
setup_mcp_servers,
|
|
teardown_mcp_servers,
|
|
)
|
|
|
|
|
|
def _stdio(toolset: MCPToolset) -> StdioTransport:
|
|
transport = toolset.client.transport
|
|
assert isinstance(transport, StdioTransport)
|
|
return transport
|
|
|
|
|
|
def _http(toolset: MCPToolset) -> StreamableHttpTransport:
|
|
transport = toolset.client.transport
|
|
assert isinstance(transport, StreamableHttpTransport)
|
|
return transport
|
|
|
|
|
|
class TestBuildToolsets:
|
|
def test_stdio_config(self):
|
|
built = _build_toolsets(
|
|
{"stdio": {"command": "mcp", "args": ["run", "x"], "env": {"FOO": "bar"}}}
|
|
)
|
|
env = _stdio(built["stdio"]).env or {}
|
|
assert env["LOGLEVEL"] == "ERROR"
|
|
assert env["FOO"] == "bar"
|
|
|
|
def test_stdio_does_not_inherit_parent_env(self, monkeypatch):
|
|
"""Secure default: parent env (PATH, secrets, etc.) is not leaked."""
|
|
monkeypatch.setenv("SECRET_TOKEN", "leak-me")
|
|
built = _build_toolsets({"stdio": {"command": "mcp", "args": []}})
|
|
env = _stdio(built["stdio"]).env or {}
|
|
assert "SECRET_TOKEN" not in env
|
|
assert "PATH" not in env
|
|
|
|
def test_stdio_user_env_overrides_logging_overrides(self):
|
|
built = _build_toolsets(
|
|
{"stdio": {"command": "mcp", "args": [], "env": {"LOGLEVEL": "DEBUG"}}}
|
|
)
|
|
env = _stdio(built["stdio"]).env or {}
|
|
assert env["LOGLEVEL"] == "DEBUG"
|
|
|
|
def test_stdio_cwd_is_forwarded(self, tmp_path):
|
|
built = _build_toolsets(
|
|
{"stdio": {"command": "mcp", "args": [], "cwd": str(tmp_path)}}
|
|
)
|
|
assert _stdio(built["stdio"]).cwd == str(tmp_path)
|
|
|
|
def test_http_config(self):
|
|
built = _build_toolsets({"http": {"url": "http://example.com/mcp"}})
|
|
assert isinstance(built["http"].client.transport, StreamableHttpTransport)
|
|
|
|
def test_http_with_authorization_header(self):
|
|
built = _build_toolsets(
|
|
{
|
|
"http": {
|
|
"url": "http://example.com/mcp",
|
|
"headers": {"Authorization": "Bearer secret"},
|
|
}
|
|
}
|
|
)
|
|
transport = _http(built["http"])
|
|
assert transport.headers == {"Authorization": "Bearer secret"}
|
|
|
|
def test_url_ending_in_sse_resolves_to_sse_transport(self):
|
|
built = _build_toolsets({"sse": {"url": "http://example.com/sse"}})
|
|
assert isinstance(built["sse"].client.transport, SSETransport)
|
|
|
|
def test_websocket_url_rejected(self):
|
|
with pytest.raises(ValueError, match="WebSocket transport"):
|
|
_build_toolsets({"ws": {"url": "ws://example.com/mcp"}})
|
|
|
|
def test_wss_url_rejected(self):
|
|
with pytest.raises(ValueError, match="WebSocket transport"):
|
|
_build_toolsets({"wss": {"url": "wss://example.com/mcp"}})
|
|
|
|
@pytest.mark.parametrize("name", ["k8s.lab", "my server", "a/b", "grafana!"])
|
|
def test_server_name_with_provider_invalid_characters_rejected(self, name):
|
|
"""The name prefixes every tool sent to the model, so providers see it."""
|
|
with pytest.raises(ValueError, match="letters, digits"):
|
|
_build_toolsets({name: {"url": "http://example.com/mcp"}})
|
|
|
|
@pytest.mark.parametrize("name", ["k8s", "my-server", "my_server", "grafana2"])
|
|
def test_provider_valid_server_names_accepted(self, name):
|
|
built = _build_toolsets({name: {"url": "http://example.com/mcp"}})
|
|
assert name in built
|
|
|
|
def test_env_var_substitution_in_env_values(self, monkeypatch):
|
|
monkeypatch.setenv("MY_TOKEN", "shh")
|
|
built = _build_toolsets(
|
|
{
|
|
"stdio": {
|
|
"command": "mcp",
|
|
"args": [],
|
|
"env": {"GITHUB_TOKEN": "${MY_TOKEN}"},
|
|
}
|
|
}
|
|
)
|
|
env = _stdio(built["stdio"]).env or {}
|
|
assert env["GITHUB_TOKEN"] == "shh"
|
|
|
|
def test_env_var_substitution_in_command_and_args(self, monkeypatch):
|
|
monkeypatch.setenv("MCP_BIN", "/opt/bin/mcp")
|
|
built = _build_toolsets(
|
|
{
|
|
"stdio": {
|
|
"command": "${MCP_BIN}",
|
|
"args": ["--config", "${MCP_BIN}.conf"],
|
|
}
|
|
}
|
|
)
|
|
transport = _stdio(built["stdio"])
|
|
assert transport.command == "/opt/bin/mcp"
|
|
assert list(transport.args) == ["--config", "/opt/bin/mcp.conf"]
|
|
|
|
def test_env_var_substitution_with_default(self, monkeypatch):
|
|
monkeypatch.delenv("MISSING_VAR", raising=False)
|
|
built = _build_toolsets(
|
|
{
|
|
"stdio": {
|
|
"command": "mcp",
|
|
"args": [],
|
|
"env": {"DEFAULTED": "${MISSING_VAR:-fallback}"},
|
|
}
|
|
}
|
|
)
|
|
env = _stdio(built["stdio"]).env or {}
|
|
assert env["DEFAULTED"] == "fallback"
|
|
|
|
def test_env_var_substitution_missing_raises(self, monkeypatch):
|
|
monkeypatch.delenv("UNDEFINED_VAR", raising=False)
|
|
with pytest.raises(ValueError, match="UNDEFINED_VAR"):
|
|
_build_toolsets(
|
|
{
|
|
"stdio": {
|
|
"command": "mcp",
|
|
"args": [],
|
|
"env": {"X": "${UNDEFINED_VAR}"},
|
|
}
|
|
}
|
|
)
|
|
|
|
def test_env_var_substitution_in_authorization_header(self, monkeypatch):
|
|
monkeypatch.setenv("BEARER", "s3cret")
|
|
built = _build_toolsets(
|
|
{
|
|
"http": {
|
|
"url": "http://x/mcp",
|
|
"headers": {"Authorization": "Bearer ${BEARER}"},
|
|
}
|
|
}
|
|
)
|
|
transport = _http(built["http"])
|
|
assert transport.headers == {"Authorization": "Bearer s3cret"}
|
|
|
|
|
|
class TestSetupAndTeardown:
|
|
async def test_no_config_returns_empty(self, app_config):
|
|
meta = await setup_mcp_servers()
|
|
assert meta == {}
|
|
assert mcp_servers == {}
|
|
await teardown_mcp_servers()
|
|
|
|
async def test_stdio_server_loads_tools(self, app_config, mcp_server_config):
|
|
app_config.set("mcpServers", {"test_server": mcp_server_config["stdio"]})
|
|
try:
|
|
meta = await setup_mcp_servers()
|
|
assert "test_server" in meta
|
|
names = {m["name"] for m in meta["test_server"]}
|
|
assert {"oracle", "puzzle_solver"}.issubset(names)
|
|
assert "test_server" in mcp_servers
|
|
finally:
|
|
await teardown_mcp_servers()
|
|
|
|
async def test_failed_server_init_is_skipped(self, app_config):
|
|
app_config.set(
|
|
"mcpServers",
|
|
{"broken": {"command": "nonexistent-command-xyz", "args": []}},
|
|
)
|
|
try:
|
|
meta = await setup_mcp_servers()
|
|
assert meta == {}
|
|
assert mcp_servers == {}
|
|
finally:
|
|
await teardown_mcp_servers()
|
|
|
|
async def test_websocket_in_config_logs_and_skips_all(self, app_config):
|
|
"""Validation failures for the whole config block are logged once."""
|
|
import oterm.log
|
|
|
|
app_config.set(
|
|
"mcpServers",
|
|
{"ws-bad": {"url": "ws://localhost/mcp"}},
|
|
)
|
|
before = len(oterm.log.log_lines)
|
|
try:
|
|
meta = await setup_mcp_servers()
|
|
assert meta == {}
|
|
messages = [msg for _, msg in oterm.log.log_lines[before:]]
|
|
assert any("WebSocket" in m for m in messages)
|
|
finally:
|
|
await teardown_mcp_servers()
|
|
|
|
async def test_invalid_server_name_in_config_logs_and_skips_all(self, app_config):
|
|
import oterm.log
|
|
|
|
app_config.set("mcpServers", {"k8s.lab": {"url": "http://localhost/mcp"}})
|
|
before = len(oterm.log.log_lines)
|
|
try:
|
|
meta = await setup_mcp_servers()
|
|
assert meta == {}
|
|
messages = [msg for _, msg in oterm.log.log_lines[before:]]
|
|
assert any("k8s.lab" in m and "letters, digits" in m for m in messages)
|
|
finally:
|
|
await teardown_mcp_servers()
|
|
|
|
async def test_tool_over_provider_name_limit_is_dropped(
|
|
self, app_config, mcp_server_config
|
|
):
|
|
"""`{server}_{tool}` must stay within the 64 characters providers allow."""
|
|
import oterm.log
|
|
|
|
long_name = "s" * 60
|
|
app_config.set("mcpServers", {long_name: mcp_server_config["stdio"]})
|
|
before = len(oterm.log.log_lines)
|
|
try:
|
|
meta = await setup_mcp_servers()
|
|
assert meta[long_name] == []
|
|
messages = [msg for _, msg in oterm.log.log_lines[before:]]
|
|
assert any("oracle" in m and "64" in m for m in messages)
|
|
finally:
|
|
await teardown_mcp_servers()
|
|
|
|
async def test_tools_within_provider_name_limit_are_kept(
|
|
self, app_config, mcp_server_config
|
|
):
|
|
app_config.set("mcpServers", {"short": mcp_server_config["stdio"]})
|
|
try:
|
|
meta = await setup_mcp_servers()
|
|
assert {m["name"] for m in meta["short"]} == {"oracle", "puzzle_solver"}
|
|
finally:
|
|
await teardown_mcp_servers()
|
|
|
|
async def test_unexpected_exception_during_build_is_logged(
|
|
self, app_config, monkeypatch
|
|
):
|
|
"""Non-ValueError exceptions during _build_toolsets are caught and logged."""
|
|
import oterm.log
|
|
import oterm.tools.mcp.setup as setup_mod
|
|
|
|
def boom(_raw):
|
|
raise RuntimeError("kaboom")
|
|
|
|
monkeypatch.setattr(setup_mod, "_build_toolsets", boom)
|
|
app_config.set("mcpServers", {"x": {"url": "http://example.com/mcp"}})
|
|
before = len(oterm.log.log_lines)
|
|
try:
|
|
meta = await setup_mcp_servers()
|
|
assert meta == {}
|
|
messages = [msg for _, msg in oterm.log.log_lines[before:]]
|
|
assert any("could not be parsed" in m for m in messages)
|
|
finally:
|
|
await teardown_mcp_servers()
|