|
|
@@ -0,0 +1,120 @@
|
|
|
+from __future__ import annotations
|
|
|
+
|
|
|
+import asyncio
|
|
|
+import json
|
|
|
+import logging
|
|
|
+import unittest
|
|
|
+from unittest.mock import patch
|
|
|
+
|
|
|
+from fastmcp.server.middleware import MiddlewareContext
|
|
|
+from mcp.types import CallToolRequestParams
|
|
|
+
|
|
|
+from data_collector_mcp.tool_logging import ToolCallLoggingMiddleware
|
|
|
+
|
|
|
+
|
|
|
+class ToolCallLoggingMiddlewareTests(unittest.TestCase):
|
|
|
+ def test_success_logs_start_and_end_with_same_trace_id(self) -> None:
|
|
|
+ middleware = ToolCallLoggingMiddleware()
|
|
|
+ context = MiddlewareContext(
|
|
|
+ message=CallToolRequestParams(
|
|
|
+ name="collector.device_list",
|
|
|
+ arguments={"project_key": "dev-01", "num_points": False},
|
|
|
+ ),
|
|
|
+ method="tools/call",
|
|
|
+ )
|
|
|
+
|
|
|
+ async def call_next(received_context: MiddlewareContext) -> dict[str, int]:
|
|
|
+ self.assertIs(received_context, context)
|
|
|
+ return {"state": 0}
|
|
|
+
|
|
|
+ with self.assertLogs("data_collector_mcp.tool_calls", level="INFO") as logs:
|
|
|
+ result = asyncio.run(middleware.on_call_tool(context, call_next))
|
|
|
+
|
|
|
+ self.assertEqual(result, {"state": 0})
|
|
|
+ payloads = [json.loads(line.split("INFO:data_collector_mcp.tool_calls:", 1)[1]) for line in logs.output]
|
|
|
+ self.assertEqual([item["event"] for item in payloads], ["mcp_tool_start", "mcp_tool_end"])
|
|
|
+ self.assertEqual(payloads[0]["trace_id"], payloads[1]["trace_id"])
|
|
|
+ self.assertEqual(payloads[0]["tool_name"], "collector.device_list")
|
|
|
+ self.assertEqual(payloads[0]["arguments"], {"project_key": "dev-01", "num_points": False})
|
|
|
+ self.assertFalse(payloads[0]["arguments_truncated"])
|
|
|
+ self.assertEqual(payloads[1]["status"], "success")
|
|
|
+ self.assertIn("duration_ms", payloads[1])
|
|
|
+
|
|
|
+ def test_each_call_gets_unique_trace_id(self) -> None:
|
|
|
+ middleware = ToolCallLoggingMiddleware()
|
|
|
+ context = MiddlewareContext(
|
|
|
+ message=CallToolRequestParams(name="project.list", arguments={}),
|
|
|
+ method="tools/call",
|
|
|
+ )
|
|
|
+
|
|
|
+ async def call_next(received_context: MiddlewareContext) -> dict[str, int]:
|
|
|
+ return {"total": 0}
|
|
|
+
|
|
|
+ with self.assertLogs("data_collector_mcp.tool_calls", level="INFO") as logs:
|
|
|
+ asyncio.run(middleware.on_call_tool(context, call_next))
|
|
|
+ asyncio.run(middleware.on_call_tool(context, call_next))
|
|
|
+
|
|
|
+ starts = [
|
|
|
+ json.loads(line.split("INFO:data_collector_mcp.tool_calls:", 1)[1])
|
|
|
+ for line in logs.output
|
|
|
+ if "mcp_tool_start" in line
|
|
|
+ ]
|
|
|
+ self.assertEqual(len(starts), 2)
|
|
|
+ self.assertNotEqual(starts[0]["trace_id"], starts[1]["trace_id"])
|
|
|
+
|
|
|
+ def test_error_logs_start_and_error_with_same_trace_id(self) -> None:
|
|
|
+ middleware = ToolCallLoggingMiddleware()
|
|
|
+ context = MiddlewareContext(
|
|
|
+ message=CallToolRequestParams(name="collector.device_list", arguments={"project_key": "missing"}),
|
|
|
+ method="tools/call",
|
|
|
+ )
|
|
|
+
|
|
|
+ async def call_next(received_context: MiddlewareContext) -> dict[str, int]:
|
|
|
+ raise ValueError("project_key not found: missing")
|
|
|
+
|
|
|
+ with self.assertLogs("data_collector_mcp.tool_calls", level="INFO") as logs:
|
|
|
+ with self.assertRaisesRegex(ValueError, "project_key not found"):
|
|
|
+ asyncio.run(middleware.on_call_tool(context, call_next))
|
|
|
+
|
|
|
+ start_payload = json.loads(logs.output[0].split("INFO:data_collector_mcp.tool_calls:", 1)[1])
|
|
|
+ error_payload = json.loads(logs.output[1].split("ERROR:data_collector_mcp.tool_calls:", 1)[1])
|
|
|
+ self.assertEqual(start_payload["event"], "mcp_tool_start")
|
|
|
+ self.assertEqual(error_payload["event"], "mcp_tool_error")
|
|
|
+ self.assertEqual(start_payload["trace_id"], error_payload["trace_id"])
|
|
|
+ self.assertEqual(error_payload["status"], "error")
|
|
|
+ self.assertEqual(error_payload["error_type"], "ValueError")
|
|
|
+ self.assertEqual(error_payload["error"], "project_key not found: missing")
|
|
|
+
|
|
|
+ def test_arguments_are_redacted_and_truncated(self) -> None:
|
|
|
+ middleware = ToolCallLoggingMiddleware()
|
|
|
+ context = MiddlewareContext(
|
|
|
+ message=CallToolRequestParams(
|
|
|
+ name="example.secret_tool",
|
|
|
+ arguments={
|
|
|
+ "project_key": "dev-01",
|
|
|
+ "password": "plain-text",
|
|
|
+ "nested": {"token": "abc123"},
|
|
|
+ "large": "x" * 100,
|
|
|
+ },
|
|
|
+ ),
|
|
|
+ method="tools/call",
|
|
|
+ )
|
|
|
+
|
|
|
+ async def call_next(received_context: MiddlewareContext) -> dict[str, int]:
|
|
|
+ return {"state": 0}
|
|
|
+
|
|
|
+ with patch.dict("os.environ", {"MCP_LOG_ARGUMENT_MAX_LENGTH": "60"}):
|
|
|
+ with self.assertLogs("data_collector_mcp.tool_calls", level="INFO") as logs:
|
|
|
+ asyncio.run(middleware.on_call_tool(context, call_next))
|
|
|
+
|
|
|
+ start_payload = json.loads(logs.output[0].split("INFO:data_collector_mcp.tool_calls:", 1)[1])
|
|
|
+ self.assertTrue(start_payload["arguments_truncated"])
|
|
|
+ self.assertIsInstance(start_payload["arguments"], str)
|
|
|
+ self.assertIn("***REDACTED***", start_payload["arguments"])
|
|
|
+ self.assertNotIn("plain-text", start_payload["arguments"])
|
|
|
+ self.assertNotIn("abc123", start_payload["arguments"])
|
|
|
+
|
|
|
+
|
|
|
+if __name__ == "__main__":
|
|
|
+ logging.basicConfig(level=logging.INFO)
|
|
|
+ unittest.main()
|