This is an automated email from the ASF dual-hosted git repository.
rusackas pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/superset.git
The following commit(s) were added to refs/heads/master by this push:
new 181f5bc9c1e fix(mcp): capture created chart/dashboard id in
generate_chart/generate_dashboard audit logs (#42367)
181f5bc9c1e is described below
commit 181f5bc9c1ed38b9ee4b589b907fdf11f5f5f22e
Author: JUST.in DO IT <[email protected]>
AuthorDate: Tue Jul 28 09:16:58 2026 -0700
fix(mcp): capture created chart/dashboard id in
generate_chart/generate_dashboard audit logs (#42367)
---
superset/mcp_service/middleware.py | 222 ++++++++++++++++-----
.../mcp_service/test_middleware_logging.py | 130 ++++++++++++
2 files changed, 305 insertions(+), 47 deletions(-)
diff --git a/superset/mcp_service/middleware.py
b/superset/mcp_service/middleware.py
index 5faebbb19c7..363a8e8c6ad 100644
--- a/superset/mcp_service/middleware.py
+++ b/superset/mcp_service/middleware.py
@@ -261,6 +261,42 @@ class LoggingMiddleware(Middleware):
dataset_id = params.get("dataset_id")
return agent_id, user_id, dashboard_id, slice_id, dataset_id, params
+ def _extract_output_ids(self, result: ToolResult) -> tuple[int | None, int
| None]:
+ """Extract dashboard/chart IDs created by the tool from its response.
+
+ Create-style tools (generate_chart, generate_dashboard) don't take
+ chart_id/dashboard_id as input, so _extract_context_info never sees
+ them and every retry logs slice_id/dashboard_id=None even on the
+ attempt that actually persisted the object. Look at the response
+ body instead, since that's the only place the new ID appears.
+ Supports both flat ("chart_id"/"dashboard_id") and nested
+ ("chart"/"dashboard" objects with an "id" field) response shapes.
+ """
+ from superset.utils.json import loads as json_loads
+
+ try:
+ data = json_loads(result.content[0].text)
+ except (AttributeError, IndexError, ValueError, TypeError):
+ return None, None
+ if not isinstance(data, dict):
+ return None, None
+
+ slice_id = None
+ chart = data.get("chart")
+ if isinstance(chart, dict):
+ slice_id = chart.get("id")
+ if slice_id is None:
+ slice_id = data.get("chart_id")
+
+ dashboard_id = None
+ dashboard = data.get("dashboard")
+ if isinstance(dashboard, dict):
+ dashboard_id = dashboard.get("id")
+ if dashboard_id is None:
+ dashboard_id = data.get("dashboard_id")
+
+ return dashboard_id, slice_id
+
@staticmethod
def _resolve_tool_name(tool_name: str | None, params: Any) -> str | None:
"""Resolve the underlying tool name from call_tool proxy arguments.
@@ -283,6 +319,129 @@ class LoggingMiddleware(Middleware):
return params["name"]
return None
+ def _backfill_output_ids(
+ self,
+ success: bool,
+ result: Any,
+ dashboard_id: int | None,
+ slice_id: int | None,
+ ) -> tuple[int | None, int | None]:
+ """Fill in missing ids from a create tool's response on success.
+
+ Create-style tools (generate_chart, generate_dashboard) don't take
+ the new object's ID as input, so it's missing from params. On a
+ successful call, pull it from the response instead so retried
+ creates are distinguishable.
+ """
+ if not success or not isinstance(result, ToolResult):
+ return dashboard_id, slice_id
+ output_dashboard_id, output_slice_id = self._extract_output_ids(result)
+ if dashboard_id is None:
+ dashboard_id = output_dashboard_id
+ if slice_id is None:
+ slice_id = output_slice_id
+ return dashboard_id, slice_id
+
+ @staticmethod
+ def _build_call_tool_payload(
+ *,
+ mcp_call_id: str,
+ tool_name: str | None,
+ agent_id: str | None,
+ params: Any,
+ method: str,
+ dashboard_id: int | None,
+ slice_id: int | None,
+ dataset_id: int | None,
+ success: bool,
+ mcp_tool: str | None,
+ error_type: str | None,
+ ) -> dict[str, Any]:
+ payload: dict[str, Any] = {
+ "mcp_call_id": mcp_call_id,
+ "tool": tool_name,
+ "agent_id": agent_id,
+ "params": _sanitize_params(params),
+ "method": method,
+ "dashboard_id": dashboard_id,
+ "slice_id": slice_id,
+ "dataset_id": dataset_id,
+ "success": success,
+ }
+ if mcp_tool is not None:
+ payload["mcp_tool"] = mcp_tool
+ if error_type is not None:
+ payload["error_type"] = error_type
+ return payload
+
+ def _log_call_tool_result(
+ self,
+ *,
+ context: MiddlewareContext,
+ tool_name: str | None,
+ mcp_tool: str | None,
+ mcp_call_id: str,
+ agent_id: str | None,
+ user_id: int | None,
+ dashboard_id: int | None,
+ slice_id: int | None,
+ dataset_id: int | None,
+ params: Any,
+ success: bool,
+ error_type: str | None,
+ result: Any,
+ start_time: float,
+ ) -> None:
+ duration_ms = int((time.time() - start_time) * 1000)
+ dashboard_id, slice_id = self._backfill_output_ids(
+ success, result, dashboard_id, slice_id
+ )
+ payload = self._build_call_tool_payload(
+ mcp_call_id=mcp_call_id,
+ tool_name=tool_name,
+ agent_id=agent_id,
+ params=params,
+ method=context.method,
+ dashboard_id=dashboard_id,
+ slice_id=slice_id,
+ dataset_id=dataset_id,
+ success=success,
+ mcp_tool=mcp_tool,
+ error_type=error_type,
+ )
+ if has_app_context():
+ event_logger.log(
+ user_id=user_id,
+ action="mcp_tool_call",
+ dashboard_id=dashboard_id,
+ duration_ms=duration_ms,
+ slice_id=slice_id,
+ referrer=None,
+ curated_payload=payload,
+ )
+ extra_parts = []
+ if mcp_tool is not None:
+ extra_parts.append(f"mcp_tool={mcp_tool}")
+ if error_type is not None:
+ extra_parts.append(f"error_type={error_type}")
+ extra = (", " + ", ".join(extra_parts)) if extra_parts else ""
+ logger.info(
+ "MCP tool call: tool=%s, agent_id=%s, user_id=%s, method=%s, "
+ "dashboard_id=%s, slice_id=%s, dataset_id=%s, duration_ms=%s, "
+ "success=%s, mcp_call_id=%s%s",
+ tool_name,
+ agent_id,
+ user_id,
+ context.method,
+ dashboard_id,
+ slice_id,
+ dataset_id,
+ duration_ms,
+ success,
+ mcp_call_id,
+ extra,
+ )
+
async def on_call_tool(
self,
context: MiddlewareContext,
@@ -300,6 +459,7 @@ class LoggingMiddleware(Middleware):
start_time = time.time()
success = False
error_type: str | None = None
+ result: Any = None
try:
result = await call_next(context)
success = not self._is_error_response(result)
@@ -316,53 +476,21 @@ class LoggingMiddleware(Middleware):
success = False
raise
finally:
- duration_ms = int((time.time() - start_time) * 1000)
- payload: dict[str, Any] = {
- "mcp_call_id": mcp_call_id,
- "tool": tool_name,
- "agent_id": agent_id,
- "params": _sanitize_params(params),
- "method": context.method,
- "dashboard_id": dashboard_id,
- "slice_id": slice_id,
- "dataset_id": dataset_id,
- "success": success,
- }
- if mcp_tool is not None:
- payload["mcp_tool"] = mcp_tool
- if error_type is not None:
- payload["error_type"] = error_type
- if has_app_context():
- event_logger.log(
- user_id=user_id,
- action="mcp_tool_call",
- dashboard_id=dashboard_id,
- duration_ms=duration_ms,
- slice_id=slice_id,
- referrer=None,
- curated_payload=payload,
- )
- extra_parts = []
- if mcp_tool is not None:
- extra_parts.append(f"mcp_tool={mcp_tool}")
- if error_type is not None:
- extra_parts.append(f"error_type={error_type}")
- extra = (", " + ", ".join(extra_parts)) if extra_parts else ""
- logger.info(
- "MCP tool call: tool=%s, agent_id=%s, user_id=%s, method=%s, "
- "dashboard_id=%s, slice_id=%s, dataset_id=%s, duration_ms=%s, "
- "success=%s, mcp_call_id=%s%s",
- tool_name,
- agent_id,
- user_id,
- context.method,
- dashboard_id,
- slice_id,
- dataset_id,
- duration_ms,
- success,
- mcp_call_id,
- extra,
+ self._log_call_tool_result(
+ context=context,
+ tool_name=tool_name,
+ mcp_tool=mcp_tool,
+ mcp_call_id=mcp_call_id,
+ agent_id=agent_id,
+ user_id=user_id,
+ dashboard_id=dashboard_id,
+ slice_id=slice_id,
+ dataset_id=dataset_id,
+ params=params,
+ success=success,
+ error_type=error_type,
+ result=result,
+ start_time=start_time,
)
async def on_message(
diff --git a/tests/unit_tests/mcp_service/test_middleware_logging.py
b/tests/unit_tests/mcp_service/test_middleware_logging.py
index de8f246b599..d838a478a1d 100644
--- a/tests/unit_tests/mcp_service/test_middleware_logging.py
+++ b/tests/unit_tests/mcp_service/test_middleware_logging.py
@@ -216,6 +216,102 @@ class TestLoggingMiddlewareOnCallTool:
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
+
class TestLoggingMiddlewareOnMessage:
"""Tests for LoggingMiddleware.on_message()."""
@@ -468,6 +564,40 @@ class TestIsErrorResponse:
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)."""
+ middleware = LoggingMiddleware()
+ result = ToolResult(
+ content=[mt.TextContent(type="text", text="not valid json {{{")]
+ )
+ assert middleware._extract_output_ids(result) == (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(ToolResult(content=[])) == (
+ 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(result) == (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(result) == (456, 123)
+
+
class TestMiddlewareChainOrder:
"""Test that the middleware order from server.py logs failures correctly.