from __future__ import annotations import json import logging import os import time import uuid from datetime import datetime, timezone from typing import Any from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext SENSITIVE_KEYS = { "access_token", "auth_token", "authorization", "password", "refresh_token", "secret", "token", } def _utc_now_iso() -> str: return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") def _argument_max_length() -> int: try: return max(0, int(str(os.getenv("MCP_LOG_ARGUMENT_MAX_LENGTH", "10000")).strip())) except (TypeError, ValueError): return 10000 def _redact(value: Any) -> Any: if isinstance(value, dict): return { key: "***REDACTED***" if str(key).lower() in SENSITIVE_KEYS else _redact(item) for key, item in value.items() } if isinstance(value, list): return [_redact(item) for item in value] if isinstance(value, tuple): return [_redact(item) for item in value] return value def _serialize_arguments(arguments: Any) -> tuple[Any, bool]: redacted = _redact(arguments or {}) max_length = _argument_max_length() text = json.dumps(redacted, ensure_ascii=False, default=str, separators=(",", ":")) if max_length and len(text) > max_length: return text[:max_length] + "...", True return redacted, False class ToolCallLoggingMiddleware(Middleware): def __init__(self, logger: logging.Logger | None = None) -> None: self.logger = logger or logging.getLogger("data_collector_mcp.tool_calls") def _log(self, payload: dict[str, Any], level: int = logging.INFO) -> None: self.logger.log(level, json.dumps(payload, ensure_ascii=False, default=str, separators=(",", ":"))) async def on_call_tool(self, context: MiddlewareContext, call_next: CallNext) -> Any: trace_id = uuid.uuid4().hex tool_name = str(getattr(context.message, "name", "unknown") or "unknown") arguments, arguments_truncated = _serialize_arguments(getattr(context.message, "arguments", None)) started_at = _utc_now_iso() start_perf = time.perf_counter() start_payload: dict[str, Any] = { "event": "mcp_tool_start", "trace_id": trace_id, "tool_name": tool_name, "arguments": arguments, "arguments_truncated": arguments_truncated, "started_at": started_at, } self._log(start_payload) try: result = await call_next(context) except Exception as exc: ended_at = _utc_now_iso() self._log( { "event": "mcp_tool_error", "trace_id": trace_id, "tool_name": tool_name, "status": "error", "started_at": started_at, "ended_at": ended_at, "duration_ms": round((time.perf_counter() - start_perf) * 1000, 2), "error_type": type(exc).__name__, "error": str(exc), }, logging.ERROR, ) raise ended_at = _utc_now_iso() self._log( { "event": "mcp_tool_end", "trace_id": trace_id, "tool_name": tool_name, "status": "success", "started_at": started_at, "ended_at": ended_at, "duration_ms": round((time.perf_counter() - start_perf) * 1000, 2), } ) return result