Files
superset2/tests/unit_tests/mcp_service/test_middleware_logging.py
T
justinpark 309aec3dab fix(mcp): resolve real tool name and capture entity IDs through call_tool proxy and request wrapper
_resolve_tool_name() was dead code (never invoked, and _CALL_TOOL_PROXY
was undefined), so tool search proxy calls always logged tool="call_tool".
Separately, every MCP tool takes a single `request` argument, so real
arguments arrive nested as {"request": {...}} (and one layer deeper when
routed through the call_tool proxy), which the previous flat
params.get("dashboard_id") lookup could never see through. Also extend
the create-tool output backfill to cover create_virtual_dataset's
dataset_id, which generate_chart/generate_dashboard already had for
chart/dashboard ids.
2026-08-03 17:21:24 -07:00

669 lines
26 KiB
Python

# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
"""
Unit tests for LoggingMiddleware on_call_tool() and on_message() methods.
Tests verify that:
- on_call_tool() captures duration_ms and success status
- on_message() logs non-tool messages without duration
- _extract_context_info() extracts entity IDs from params
"""
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import mcp.types as mt
import pytest
from fastmcp.tools.tool import ToolResult
from superset.mcp_service.middleware import LoggingMiddleware
def _make_context(
method: str = "tools/call",
name: str = "list_charts",
params: dict[str, Any] | None = None,
metadata: dict[str, Any] | None = None,
):
"""Create a mock MiddlewareContext."""
ctx = MagicMock()
ctx.method = method
message = MagicMock()
message.name = name
message.arguments = params or {}
ctx.message = message
if metadata is not None:
ctx.metadata = metadata
else:
ctx.metadata = None
ctx.session = None
return ctx
class TestLoggingMiddlewareOnCallTool:
"""Tests for LoggingMiddleware.on_call_tool()."""
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_logs_duration_and_success(
self, mock_get_user_id, mock_event_logger
):
"""on_call_tool records duration_ms and success=True on normal return."""
middleware = LoggingMiddleware()
ctx = _make_context(name="list_charts")
call_next = AsyncMock(return_value="tool_result")
result = await middleware.on_call_tool(ctx, call_next)
assert result == "tool_result"
call_next.assert_awaited_once_with(ctx)
# Verify event_logger.log was called with duration_ms and success
mock_event_logger.log.assert_called_once()
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["action"] == "mcp_tool_call"
assert call_kwargs["user_id"] == 42
assert isinstance(call_kwargs["duration_ms"], int)
assert call_kwargs["duration_ms"] >= 0
assert call_kwargs["curated_payload"]["success"] is True
assert call_kwargs["curated_payload"]["tool"] == "list_charts"
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_logs_failure_on_exception(
self, mock_get_user_id, mock_event_logger
):
"""on_call_tool records success=False when tool raises."""
middleware = LoggingMiddleware()
ctx = _make_context(name="execute_sql")
call_next = AsyncMock(side_effect=ValueError("boom"))
with pytest.raises(ValueError, match="boom"):
await middleware.on_call_tool(ctx, call_next)
# Verify event_logger.log was still called (in the finally block)
mock_event_logger.log.assert_called_once()
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["curated_payload"]["success"] is False
assert call_kwargs["duration_ms"] >= 0
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_logs_failure_on_tool_error(
self, mock_get_user_id, mock_event_logger
):
"""on_call_tool records success=False when GlobalErrorHandler raises ToolError.
This simulates the real middleware chain: GlobalErrorHandler catches
tool exceptions and re-raises them as ToolError. Since LoggingMiddleware
sits between GlobalErrorHandler and StructuredContentStripper, it
catches the ToolError directly.
"""
from fastmcp.exceptions import ToolError
middleware = LoggingMiddleware()
ctx = _make_context(name="get_chart_info")
call_next = AsyncMock(side_effect=ToolError("Chart 999999 not found"))
with pytest.raises(ToolError, match="Chart 999999 not found"):
await middleware.on_call_tool(ctx, call_next)
mock_event_logger.log.assert_called_once()
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["curated_payload"]["success"] is False
assert call_kwargs["curated_payload"]["tool"] == "get_chart_info"
assert call_kwargs["duration_ms"] >= 0
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_extracts_entity_ids(
self, mock_get_user_id, mock_event_logger
):
"""on_call_tool extracts dashboard_id, chart_id, dataset_id from params."""
middleware = LoggingMiddleware()
ctx = _make_context(
name="get_chart_info",
params={
"dashboard_id": 10,
"chart_id": 20,
"dataset_id": 30,
},
)
call_next = AsyncMock(return_value="ok")
await middleware.on_call_tool(ctx, call_next)
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["dashboard_id"] == 10
assert call_kwargs["slice_id"] == 20
assert call_kwargs["curated_payload"]["dataset_id"] == 30
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_extracts_chart_id_from_response(
self, mock_get_user_id, mock_event_logger
) -> None:
"""generate_chart takes no chart_id as input, so on a successful
create the new chart's ID must be pulled from the response body
instead -- otherwise every retry logs slice_id=None and a
successful attempt can't be told apart from the failed ones.
"""
middleware = LoggingMiddleware()
ctx = _make_context(name="generate_chart", params={"dataset_id": 5})
response_text = (
'{"success": true, "chart": {"id": 123, "slice_name": "My Chart"}}'
)
original_result = ToolResult(
content=[mt.TextContent(type="text", text=response_text)]
)
call_next = AsyncMock(return_value=original_result)
await middleware.on_call_tool(ctx, call_next)
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["slice_id"] == 123
assert call_kwargs["curated_payload"]["slice_id"] == 123
assert call_kwargs["curated_payload"]["success"] is True
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_extracts_dashboard_id_from_response(
self, mock_get_user_id, mock_event_logger
) -> None:
"""generate_dashboard likewise creates an ID that only appears in
the response, not the input params."""
middleware = LoggingMiddleware()
ctx = _make_context(name="generate_dashboard", params={"chart_ids": [1, 2]})
response_text = '{"success": true, "dashboard": {"id": 456}}'
original_result = ToolResult(
content=[mt.TextContent(type="text", text=response_text)]
)
call_next = AsyncMock(return_value=original_result)
await middleware.on_call_tool(ctx, call_next)
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["dashboard_id"] == 456
assert call_kwargs["curated_payload"]["dashboard_id"] == 456
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_does_not_extract_id_on_failed_response(
self, mock_get_user_id, mock_event_logger
) -> None:
"""A failed create (error schema response, no exception raised)
must not report a chart_id -- nothing was actually persisted."""
middleware = LoggingMiddleware()
ctx = _make_context(name="generate_chart", params={"dataset_id": 5})
response_text = (
'{"success": false, "chart": null, '
'"error": {"error_type": "validation_error"}}'
)
original_result = ToolResult(
content=[mt.TextContent(type="text", text=response_text)]
)
call_next = AsyncMock(return_value=original_result)
await middleware.on_call_tool(ctx, call_next)
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["curated_payload"]["success"] is False
assert call_kwargs["slice_id"] is None
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_prefers_input_slice_id_over_response(
self, mock_get_user_id, mock_event_logger
) -> None:
"""When chart_id is already known from input params (e.g.
update_chart), the response body must not override it."""
middleware = LoggingMiddleware()
ctx = _make_context(name="update_chart", params={"chart_id": 111})
response_text = '{"success": true, "chart": {"id": 999}}'
original_result = ToolResult(
content=[mt.TextContent(type="text", text=response_text)]
)
call_next = AsyncMock(return_value=original_result)
await middleware.on_call_tool(ctx, call_next)
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["slice_id"] == 111
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_resolves_name_and_dataset_id_through_proxy(
self, mock_get_user_id, mock_event_logger
) -> None:
"""End-to-end regression test for the reported bug: a
create_virtual_dataset call made through the ``call_tool`` search
proxy must log the real tool name (not "call_tool") and the
created dataset_id (backfilled from the response), instead of
tool="call_tool" and dataset_id=None."""
middleware = LoggingMiddleware()
ctx = _make_context(
name="call_tool",
params={
"name": "create_virtual_dataset",
"arguments": {
"request": {
"database_id": 108,
"schema": "jitney",
"table_name": "Airchat Monthly Usage Summary",
"sql": "SELECT 1",
}
},
},
)
response_text = '{"id": 42, "dataset_name": "Airchat Monthly Usage Summary"}'
original_result = ToolResult(
content=[mt.TextContent(type="text", text=response_text)]
)
call_next = AsyncMock(return_value=original_result)
await middleware.on_call_tool(ctx, call_next)
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["curated_payload"]["tool"] == "create_virtual_dataset"
assert call_kwargs["curated_payload"]["dataset_id"] == 42
assert call_kwargs["curated_payload"]["success"] is True
class TestLoggingMiddlewareOnMessage:
"""Tests for LoggingMiddleware.on_message()."""
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=1)
@pytest.mark.asyncio
async def test_on_message_logs_without_duration(
self, mock_get_user_id, mock_event_logger
):
"""on_message logs with action=mcp_message and duration_ms=None."""
middleware = LoggingMiddleware()
ctx = _make_context(method="resources/read", name="instance/metadata")
call_next = AsyncMock(return_value="resource_data")
result = await middleware.on_message(ctx, call_next)
assert result == "resource_data"
call_next.assert_awaited_once_with(ctx)
mock_event_logger.log.assert_called_once()
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["action"] == "mcp_message"
assert call_kwargs["duration_ms"] is None
# on_message should NOT have success field
assert "success" not in call_kwargs["curated_payload"]
class TestExtractContextInfo:
"""Tests for LoggingMiddleware._extract_context_info()."""
@patch("superset.mcp_service.middleware.get_user_id", return_value=99)
def test_extract_with_metadata_agent_id(self, mock_get_user_id):
"""Extracts agent_id from context.metadata."""
middleware = LoggingMiddleware()
ctx = _make_context(metadata={"agent_id": "agent-123"})
agent_id, user_id, dashboard_id, slice_id, dataset_id, params = (
middleware._extract_context_info(ctx)
)
assert agent_id == "agent-123"
assert user_id == 99
@patch(
"superset.mcp_service.middleware.get_user_id",
side_effect=RuntimeError("no Flask request context"),
)
def test_extract_handles_missing_user(self, mock_get_user_id):
"""Gracefully handles missing user context."""
middleware = LoggingMiddleware()
ctx = _make_context()
agent_id, user_id, dashboard_id, slice_id, dataset_id, params = (
middleware._extract_context_info(ctx)
)
assert user_id is None
@patch("superset.mcp_service.middleware.get_user_id", return_value=1)
def test_extract_slice_id_from_chart_id(self, mock_get_user_id):
"""Extracts slice_id from chart_id param (alias)."""
middleware = LoggingMiddleware()
ctx = _make_context(params={"chart_id": 55})
_, _, _, slice_id, _, _ = middleware._extract_context_info(ctx)
assert slice_id == 55
@patch("superset.mcp_service.middleware.get_user_id", return_value=1)
def test_extract_slice_id_from_slice_id(self, mock_get_user_id):
"""Extracts slice_id from slice_id param (fallback)."""
middleware = LoggingMiddleware()
ctx = _make_context(params={"slice_id": 66})
_, _, _, slice_id, _, _ = middleware._extract_context_info(ctx)
assert slice_id == 66
@patch("superset.mcp_service.middleware.get_user_id", return_value=1)
def test_extract_reads_arguments_on_real_call_tool_request_params(
self, mock_get_user_id
) -> None:
"""Regression test: the real MCP ``CallToolRequestParams`` object
exposes tool arguments as ``.arguments``, not ``.params`` -- a
``MagicMock``-based context would auto-vivify a ``.params``
attribute and hide a mismatch. Using the real SDK type here
ensures params/dashboard_id/etc. are actually populated instead
of silently logging as empty."""
middleware = LoggingMiddleware()
message = mt.CallToolRequestParams(
name="get_dashboard_info",
arguments={"dashboard_id": 7},
)
ctx = MagicMock()
ctx.message = message
ctx.metadata = None
ctx.session = None
agent_id, user_id, dashboard_id, slice_id, dataset_id, params = (
middleware._extract_context_info(ctx)
)
assert params == {"dashboard_id": 7}
assert dashboard_id == 7
@patch("superset.mcp_service.middleware.get_user_id", return_value=1)
def test_extract_unwraps_request_wrapper(self, mock_get_user_id) -> None:
"""Every MCP tool in this service takes a single ``request``
pydantic argument, so real arguments arrive as
``{"request": {"dashboard_id": ..., "chart_id": ...}}`` (e.g.
add_chart_to_existing_dashboard). Without unwrapping this layer,
dashboard_id/chart_id extraction silently returns None."""
middleware = LoggingMiddleware()
ctx = _make_context(
name="add_chart_to_existing_dashboard",
params={"request": {"dashboard_id": 10, "chart_id": 20}},
)
_, _, dashboard_id, slice_id, _, params = middleware._extract_context_info(ctx)
assert dashboard_id == 10
assert slice_id == 20
# The raw params (used for the logged payload) stay unwrapped.
assert params == {"request": {"dashboard_id": 10, "chart_id": 20}}
@patch("superset.mcp_service.middleware.get_user_id", return_value=1)
def test_extract_unwraps_call_tool_proxy_and_request(
self, mock_get_user_id
) -> None:
"""Reproduces the reported bug: when the client calls tools through
the ``call_tool`` search proxy, arguments arrive as
``{"name": ..., "arguments": {"request": {...}}}`` -- two layers of
wrapping. Both must be unwrapped to see dataset_id."""
middleware = LoggingMiddleware()
ctx = _make_context(
name="call_tool",
params={
"name": "create_virtual_dataset",
"arguments": {
"request": {
"database_id": 108,
"dataset_id": 42,
}
},
},
)
_, _, _, _, dataset_id, _ = middleware._extract_context_info(ctx)
assert dataset_id == 42
class TestResolveToolName:
"""Tests for LoggingMiddleware._resolve_tool_name()."""
def test_resolves_real_tool_name_from_call_tool_proxy(self) -> None:
resolved = LoggingMiddleware._resolve_tool_name(
"call_tool", {"name": "create_virtual_dataset", "arguments": {}}
)
assert resolved == "create_virtual_dataset"
def test_returns_none_when_not_call_tool_proxy(self) -> None:
resolved = LoggingMiddleware._resolve_tool_name(
"get_chart_info", {"identifier": 123}
)
assert resolved is None
def test_returns_none_when_name_missing(self) -> None:
resolved = LoggingMiddleware._resolve_tool_name("call_tool", {})
assert resolved is None
class TestIsErrorResponse:
"""Tests for LoggingMiddleware._is_error_response()."""
def test_detects_error_schema_response(self):
"""Detects ToolResult containing a serialized error schema
(ChartError, DashboardError, etc.) via "error_type" field."""
from fastmcp.tools.tool import ToolResult
from mcp import types as mt
middleware = LoggingMiddleware()
error_json = (
'{"error": "Chart 999 not found",'
' "error_type": "not_found",'
' "timestamp": "2026-04-09T00:00:00Z"}'
)
result = ToolResult(content=[mt.TextContent(type="text", text=error_json)])
assert middleware._is_error_response(result) is True
def test_success_response_not_detected_as_error(self):
"""Normal ToolResult is not detected as error."""
from fastmcp.tools.tool import ToolResult
from mcp import types as mt
middleware = LoggingMiddleware()
result = ToolResult(
content=[mt.TextContent(type="text", text="Successfully retrieved data")]
)
assert middleware._is_error_response(result) is False
def test_empty_content_not_detected_as_error(self):
"""ToolResult with empty content is not detected as error."""
from fastmcp.tools.tool import ToolResult
middleware = LoggingMiddleware()
assert middleware._is_error_response(ToolResult(content=[])) is False
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_on_call_tool_logs_failure_for_error_schema(
self, mock_get_user_id, mock_event_logger
):
"""on_call_tool logs success=False when tool returns an
error schema (e.g. ChartError)."""
from fastmcp.tools.tool import ToolResult
from mcp import types as mt
middleware = LoggingMiddleware()
ctx = _make_context(name="get_chart_info")
error_json = (
'{"error": "Chart 999999 not found",'
' "error_type": "not_found",'
' "timestamp": "2026-04-09T00:00:00Z"}'
)
error_result = ToolResult(
content=[mt.TextContent(type="text", text=error_json)]
)
call_next = AsyncMock(return_value=error_result)
result = await middleware.on_call_tool(ctx, call_next)
assert result == error_result
mock_event_logger.log.assert_called_once()
call_kwargs = mock_event_logger.log.call_args[1]
assert call_kwargs["curated_payload"]["success"] is False
assert call_kwargs["curated_payload"]["tool"] == "get_chart_info"
class TestExtractOutputIds:
"""Tests for LoggingMiddleware._extract_output_ids()."""
def test_returns_none_for_non_json_body(self) -> None:
"""A malformed/non-JSON response body must not raise -- the
defensive try/except should fall back to (None, None, None)."""
middleware = LoggingMiddleware()
result = ToolResult(
content=[mt.TextContent(type="text", text="not valid json {{{")]
)
assert middleware._extract_output_ids("generate_chart", result) == (
None,
None,
None,
)
def test_returns_none_for_empty_content(self) -> None:
"""A ToolResult with no content items must not raise."""
middleware = LoggingMiddleware()
assert middleware._extract_output_ids(
"generate_chart", ToolResult(content=[])
) == (
None,
None,
None,
)
def test_returns_none_for_non_dict_json_body(self) -> None:
"""A JSON body that parses but isn't an object (e.g. a bare
list) must not raise and must yield no IDs."""
middleware = LoggingMiddleware()
result = ToolResult(content=[mt.TextContent(type="text", text="[1, 2, 3]")])
assert middleware._extract_output_ids("generate_chart", result) == (
None,
None,
None,
)
def test_extracts_both_ids_from_flat_response(self) -> None:
middleware = LoggingMiddleware()
response_text = '{"chart_id": 123, "dashboard_id": 456}'
result = ToolResult(content=[mt.TextContent(type="text", text=response_text)])
assert middleware._extract_output_ids("generate_chart", result) == (
456,
123,
None,
)
def test_extracts_dataset_id_for_create_virtual_dataset(self) -> None:
"""create_virtual_dataset's response is a flat {"id": ...} with no
wrapper key, unlike chart/dashboard -- extraction must be gated on
tool_name so unrelated tools' "id" fields aren't misattributed."""
middleware = LoggingMiddleware()
response_text = '{"id": 789, "dataset_name": "my_dataset"}'
result = ToolResult(content=[mt.TextContent(type="text", text=response_text)])
assert middleware._extract_output_ids("create_virtual_dataset", result) == (
None,
None,
789,
)
def test_does_not_extract_dataset_id_for_unrelated_tool(self) -> None:
"""A generic "id" field on an unrelated tool's response must not be
misread as a dataset_id."""
middleware = LoggingMiddleware()
response_text = '{"id": 789}'
result = ToolResult(content=[mt.TextContent(type="text", text=response_text)])
assert middleware._extract_output_ids("get_chart_info", result) == (
None,
None,
None,
)
class TestMiddlewareChainOrder:
"""Test that the middleware order from server.py logs failures correctly.
If the order is wrong (StructuredContentStripper innermost),
it swallows exceptions before LoggingMiddleware can see them,
causing success=True for failures.
"""
@patch("superset.mcp_service.middleware.event_logger")
@patch("superset.mcp_service.middleware.get_user_id", return_value=42)
@pytest.mark.asyncio
async def test_real_middleware_chain_logs_exception_as_failure(
self, mock_get_user_id, mock_event_logger
):
"""Tool exception is logged as success=False through the
real middleware chain from build_middleware_list()."""
from functools import partial
from fastmcp.tools.tool import ToolResult
from superset.mcp_service.server import build_middleware_list
middleware_list = build_middleware_list()
async def failing_tool(context: Any) -> Any:
raise ValueError("chart not found")
# Build chain same way FastMCP does
chain = failing_tool
for mw in reversed(middleware_list):
chain = partial(mw, call_next=chain)
ctx = _make_context(name="get_chart_info")
result = await chain(ctx)
# StructuredContentStripper (outermost) must catch the re-raised
# exception and convert it to a safe ToolResult with "Error:" text.
# If it's not outermost, the exception would leak to the MCP SDK.
assert isinstance(result, ToolResult)
assert result.content[0].text.startswith("Error:")
# LoggingMiddleware must log
# success=False. If the middleware order is wrong
# (StructuredContentStripper innermost), this would be
# success=True because the exception gets swallowed
# before LoggingMiddleware sees it.
log_calls = [
c
for c in mock_event_logger.log.call_args_list
if c[1].get("action") == "mcp_tool_call"
]
assert len(log_calls) == 1
assert log_calls[0][1]["curated_payload"]["success"] is False, (
"Middleware order is wrong: StructuredContentStripper is "
"swallowing exceptions before LoggingMiddleware can detect "
"them. Ensure StructuredContentStripper is outermost "
"(first added) in build_middleware_list()."
)