Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 108 additions & 0 deletions src/security.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
"""Input-validation safeguards for OpenAPI-derived MCP tools.

FastMCP interpolates tool arguments into backend URL path templates (e.g.
``/api/v1/catalog/entities/{tagOrId}/custom-data``). A path-parameter value that
contains a separator or dot-segment can break out of its URL segment and route
the request outside the intended ``/api/v1`` prefix — a path traversal that lets
a caller probe undocumented backend routes. These guards reject such values for
the parameters that are interpolated into the request path.
"""
from __future__ import annotations

import re
from typing import Any

from fastmcp.exceptions import ToolError
from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext

from .utils.logging import get_logger

logger = get_logger(__name__)

# Sequences that let a path parameter escape its URL segment: raw separators and
# dot-segments plus their percent- and backslash-encodings. Matched
# case-insensitively against the raw (pre-substitution) argument value, so
# encoded variants are caught before the backend can decode them.
_TRAVERSAL_TOKENS: tuple[str, ...] = (
"/",
"\\",
"..",
";",
"%2f", # /
"%5c", # \
"%2e", # .
)


def contains_path_traversal(value: str) -> bool:
"""Return True if ``value`` is unsafe to interpolate into a URL path segment."""
lowered = value.lower()
if any(token in lowered for token in _TRAVERSAL_TOKENS):
return True
# Control characters (including NUL) can truncate or confuse routing.
return any(ord(char) < 0x20 for char in value)


def extract_path_param_names(openapi_spec: dict[str, Any]) -> set[str]:
"""Collect every parameter name that appears as a URL path parameter.

Only path parameters are interpolated into the request path, so they are the
values that must stay free of separators. Query and body parameters (e.g. a
free-text ``context``) are intentionally excluded.
"""
names: set[str] = set()
for path, path_item in openapi_spec.get("paths", {}).items():
# Placeholders in the path template are authoritative path parameters.
names.update(re.findall(r"\{([^}]+)\}", path))
if not isinstance(path_item, dict):
continue
# Explicit `in: path` declarations at the path or operation level.
parameter_lists = [path_item.get("parameters")]
parameter_lists += [
operation.get("parameters")
for operation in path_item.values()
if isinstance(operation, dict)
]
for parameters in parameter_lists:
for parameter in parameters or []:
if (
isinstance(parameter, dict)
and parameter.get("in") == "path"
and parameter.get("name")
):
names.add(parameter["name"])
return names


class PathTraversalGuardMiddleware(Middleware):
"""Reject tool calls whose path-parameter arguments contain traversal sequences.

Runs before the request is forwarded to the backend, so a malicious value is
never composed into a backend URL.
"""

def __init__(self, path_param_names: set[str]):
self._path_param_names = set(path_param_names)

async def on_call_tool(
self,
context: MiddlewareContext,
call_next: CallNext,
) -> Any:
arguments = context.message.arguments or {}
for name, value in arguments.items():
if (
name in self._path_param_names
and isinstance(value, str)
and contains_path_traversal(value)
):
logger.warning(
"Blocked tool call %r: path parameter %r contained an illegal sequence",
context.message.name,
name,
)
raise ToolError(
f"Invalid value for '{name}': path identifiers may not contain "
"'/', '\\', '..', ';', or their encoded forms."
)
return await call_next(context)
7 changes: 7 additions & 0 deletions src/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from .components.customizers import customize_components
from .config import Config
from .routes.mappers import custom_route_mapper
from .security import PathTraversalGuardMiddleware, extract_path_param_names
from .utils.logging import setup_logging
from .utils.openapi_resolver import resolve_refs

Expand Down Expand Up @@ -54,6 +55,12 @@ def create_mcp_server() -> FastMCP:
mcp_component_fn=customize_components,
)

# Reject path-parameter values that could traverse outside the intended
# backend route before the request is forwarded (see src/security.py).
path_param_names = extract_path_param_names(openapi_spec)
mcp_server.add_middleware(PathTraversalGuardMiddleware(path_param_names))
logger.info(f"Path-traversal guard enabled for {len(path_param_names)} path parameters")

logger.info(f"MCP server '{Config.APP_NAME}' created successfully")

return mcp_server
Expand Down
117 changes: 117 additions & 0 deletions tests/test_security.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,117 @@
"""Tests for path-traversal input-validation safeguards."""
from types import SimpleNamespace

import pytest
from fastmcp.exceptions import ToolError

from src.security import (
PathTraversalGuardMiddleware,
contains_path_traversal,
extract_path_param_names,
)


class TestContainsPathTraversal:
"""Values that must be rejected mirror the payloads in the customer report."""

@pytest.mark.parametrize(
"value",
[
"../../v1/catalog/entities/services?",
"../../internal/health?",
"../../internal/v1/catalog?",
"../../nonexistent/path",
"..%5C..%5Cinternal/health",
"..%2f..%2finternal",
"%2e%2e/internal",
"a/b",
"a\\b",
"tag;matrix=1",
"with\x00null",
],
)
def test_rejects_traversal_payloads(self, value):
assert contains_path_traversal(value) is True

@pytest.mark.parametrize(
"value",
[
"test",
"my-service",
"123",
"service-tag_v2",
"Some Display Name",
"a.b.c", # dots alone are fine; only ".." is a dot-segment
],
)
def test_allows_legitimate_identifiers(self, value):
assert contains_path_traversal(value) is False


class TestExtractPathParamNames:
def test_extracts_template_placeholders_and_declared_path_params(self):
spec = {
"paths": {
"/api/v1/catalog/entities/{tagOrId}/custom-data": {
"get": {
"parameters": [
{"name": "context", "in": "query"},
{"name": "key", "in": "path"},
]
}
},
"/api/v1/dependencies/{callerTag}": {"get": {}},
}
}
assert extract_path_param_names(spec) == {"tagOrId", "key", "callerTag"}

def test_query_only_params_are_excluded(self):
spec = {
"paths": {
"/api/v1/search": {
"get": {"parameters": [{"name": "context", "in": "query"}]}
}
}
}
assert extract_path_param_names(spec) == set()


def _context(name: str, arguments: dict):
return SimpleNamespace(message=SimpleNamespace(name=name, arguments=arguments))


class TestPathTraversalGuardMiddleware:
@pytest.fixture
def middleware(self):
return PathTraversalGuardMiddleware({"tagOrId", "callerTag", "key"})

async def _call_next(self, context):
return "forwarded"

@pytest.mark.asyncio
async def test_blocks_traversal_in_path_param(self, middleware):
context = _context(
"getCustomDataForEntity",
{"tagOrId": "../../internal/health?", "context": ""},
)
with pytest.raises(ToolError):
await middleware.on_call_tool(context, self._call_next)

@pytest.mark.asyncio
async def test_allows_traversal_like_value_in_non_path_param(self, middleware):
# `context` is not a path parameter, so slashes in it are legitimate.
context = _context(
"getCustomDataForEntity",
{"tagOrId": "my-service", "context": "look in ../docs for details"},
)
assert await middleware.on_call_tool(context, self._call_next) == "forwarded"

@pytest.mark.asyncio
async def test_allows_legitimate_call(self, middleware):
context = _context("getEntityDetails", {"tagOrId": "my-service"})
assert await middleware.on_call_tool(context, self._call_next) == "forwarded"

@pytest.mark.asyncio
async def test_allows_call_with_no_arguments(self, middleware):
context = _context("listAllEntities", {})
assert await middleware.on_call_tool(context, self._call_next) == "forwarded"
Loading