tool_logging.py 3.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114
  1. from __future__ import annotations
  2. import json
  3. import logging
  4. import os
  5. import time
  6. import uuid
  7. from datetime import datetime, timezone
  8. from typing import Any
  9. from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
  10. SENSITIVE_KEYS = {
  11. "access_token",
  12. "auth_token",
  13. "authorization",
  14. "password",
  15. "refresh_token",
  16. "secret",
  17. "token",
  18. }
  19. def _utc_now_iso() -> str:
  20. return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z")
  21. def _argument_max_length() -> int:
  22. try:
  23. return max(0, int(str(os.getenv("MCP_LOG_ARGUMENT_MAX_LENGTH", "10000")).strip()))
  24. except (TypeError, ValueError):
  25. return 10000
  26. def _redact(value: Any) -> Any:
  27. if isinstance(value, dict):
  28. return {
  29. key: "***REDACTED***" if str(key).lower() in SENSITIVE_KEYS else _redact(item)
  30. for key, item in value.items()
  31. }
  32. if isinstance(value, list):
  33. return [_redact(item) for item in value]
  34. if isinstance(value, tuple):
  35. return [_redact(item) for item in value]
  36. return value
  37. def _serialize_arguments(arguments: Any) -> tuple[Any, bool]:
  38. redacted = _redact(arguments or {})
  39. max_length = _argument_max_length()
  40. text = json.dumps(redacted, ensure_ascii=False, default=str, separators=(",", ":"))
  41. if max_length and len(text) > max_length:
  42. return text[:max_length] + "...", True
  43. return redacted, False
  44. class ToolCallLoggingMiddleware(Middleware):
  45. def __init__(self, logger: logging.Logger | None = None) -> None:
  46. self.logger = logger or logging.getLogger("data_collector_mcp.tool_calls")
  47. def _log(self, payload: dict[str, Any], level: int = logging.INFO) -> None:
  48. self.logger.log(level, json.dumps(payload, ensure_ascii=False, default=str, separators=(",", ":")))
  49. async def on_call_tool(self, context: MiddlewareContext, call_next: CallNext) -> Any:
  50. trace_id = uuid.uuid4().hex
  51. tool_name = str(getattr(context.message, "name", "unknown") or "unknown")
  52. arguments, arguments_truncated = _serialize_arguments(getattr(context.message, "arguments", None))
  53. started_at = _utc_now_iso()
  54. start_perf = time.perf_counter()
  55. start_payload: dict[str, Any] = {
  56. "event": "mcp_tool_start",
  57. "trace_id": trace_id,
  58. "tool_name": tool_name,
  59. "arguments": arguments,
  60. "arguments_truncated": arguments_truncated,
  61. "started_at": started_at,
  62. }
  63. self._log(start_payload)
  64. try:
  65. result = await call_next(context)
  66. except Exception as exc:
  67. ended_at = _utc_now_iso()
  68. self._log(
  69. {
  70. "event": "mcp_tool_error",
  71. "trace_id": trace_id,
  72. "tool_name": tool_name,
  73. "status": "error",
  74. "started_at": started_at,
  75. "ended_at": ended_at,
  76. "duration_ms": round((time.perf_counter() - start_perf) * 1000, 2),
  77. "error_type": type(exc).__name__,
  78. "error": str(exc),
  79. },
  80. logging.ERROR,
  81. )
  82. raise
  83. ended_at = _utc_now_iso()
  84. self._log(
  85. {
  86. "event": "mcp_tool_end",
  87. "trace_id": trace_id,
  88. "tool_name": tool_name,
  89. "status": "success",
  90. "started_at": started_at,
  91. "ended_at": ended_at,
  92. "duration_ms": round((time.perf_counter() - start_perf) * 1000, 2),
  93. }
  94. )
  95. return result