| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114 |
- 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
|