public_server.py 9.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206
  1. import json
  2. import logging
  3. import threading
  4. import time
  5. import uuid
  6. from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
  7. from constants import DEVICE_INVALID_MESSAGE
  8. from mcp_protocol import McpProtocolHandler
  9. from services.request_context import RequestContextParser
  10. from utils.rate_limiter import SimpleRateLimiter
  11. logger = logging.getLogger(__name__)
  12. def extract_client_ip(headers, client_address):
  13. if client_address and len(client_address) > 0:
  14. return str(client_address[0])
  15. return ''
  16. class PublicMcpHttpHandler:
  17. def __init__(self, gateway_app, context_parser=None, rate_limiter=None):
  18. self.gateway_app = gateway_app
  19. self.context_parser = context_parser or RequestContextParser()
  20. self.rate_limiter = rate_limiter
  21. # Cache registered tool names at startup for rate-key validation;
  22. # unknown names fall back to the bare IP bucket, preventing bucket explosion
  23. self._known_tools = frozenset(gateway_app.registered_tool_names())
  24. def _check_rate_limit(self, rate_key, method, request_id=None, *, log_identity=None):
  25. """Returns an error response if rate limit exceeded, else None.
  26. rate_key — the bucket key used by the limiter (session-based for tools/call,
  27. IP-based for tools/list)
  28. log_identity — optional string shown in warning logs (e.g. client_ip)
  29. """
  30. if self.rate_limiter and rate_key and not self.rate_limiter.is_allowed(rate_key):
  31. logger.warning(
  32. f"[RATE_LIMIT] rate limit exceeded: identity={log_identity or rate_key}, method={method}"
  33. )
  34. return McpProtocolHandler._error_response(request_id, -32000, 'Rate limit exceeded. Please try again later.')
  35. return None
  36. def handle_json_rpc(self, headers, message, client_ip=''):
  37. request_id = message.get('id') if isinstance(message, dict) else None
  38. method = str((message or {}).get('method') or '').strip()
  39. message_params = (message or {}).get('params') or {}
  40. tool_name = str(message_params.get('name') or '').strip() \
  41. if isinstance(message_params, dict) else ''
  42. try:
  43. if method == 'initialize':
  44. # initialize is a stateless handshake that only returns server metadata;
  45. # rate-limiting it would block clients from connecting at all, so we skip it.
  46. return McpProtocolHandler._success_response(request_id, {
  47. 'protocolVersion': McpProtocolHandler.protocol_version,
  48. 'capabilities': {'tools': {'listChanged': False}},
  49. 'serverInfo': {
  50. 'name': McpProtocolHandler.server_name,
  51. 'version': McpProtocolHandler.server_version,
  52. },
  53. })
  54. if method == 'tools/list':
  55. # Keep list traffic in a dedicated IP-based bucket so it does not
  56. # compete with the per-session tools/call quota.
  57. blocked = self._check_rate_limit(
  58. 'list:{0}'.format(client_ip), method, request_id,
  59. log_identity=client_ip,
  60. )
  61. if blocked:
  62. return blocked
  63. context = self.context_parser.parse(headers or {})
  64. if not context.has_session():
  65. raise RuntimeError(DEVICE_INVALID_MESSAGE)
  66. tools = []
  67. for tool in self.gateway_app.list_tools(context.gateway_session_id):
  68. normalized = dict(tool)
  69. if 'input_schema' in normalized:
  70. normalized['inputSchema'] = normalized.pop('input_schema')
  71. tools.append(normalized)
  72. return McpProtocolHandler._success_response(request_id, {'tools': tools})
  73. if method == 'tools/call':
  74. if not isinstance(message_params, dict):
  75. raise ValueError('tool parameters must be an object')
  76. # Parse context first — tools/call always requires a valid session
  77. context = self.context_parser.parse(headers or {})
  78. if not context.has_session():
  79. raise RuntimeError(DEVICE_INVALID_MESSAGE)
  80. # Session-based rate limiting: each employee gets an independent quota per tool.
  81. # Unknown tool names fall back to the bare session bucket to prevent key explosion.
  82. session_id = context.gateway_session_id
  83. rate_key = '{0}:{1}'.format(session_id, tool_name) if tool_name in self._known_tools else session_id
  84. blocked = self._check_rate_limit(rate_key, method, request_id, log_identity=client_ip)
  85. if blocked:
  86. return blocked
  87. params = message_params
  88. result = self.gateway_app.call_tool(
  89. session_id,
  90. params.get('name'),
  91. params.get('arguments') or {},
  92. request_id='rq_http_{0}_{1}'.format(request_id, uuid.uuid4().hex[:16]),
  93. client_ip=client_ip,
  94. )
  95. return McpProtocolHandler._tool_call_response(
  96. request_id,
  97. tool_name,
  98. result,
  99. )
  100. return McpProtocolHandler._error_response(request_id, -32601, 'Method not found: {0}'.format(method))
  101. except Exception as exc:
  102. if method == 'tools/call':
  103. return McpProtocolHandler._tool_exception_response(
  104. request_id,
  105. tool_name,
  106. exc,
  107. )
  108. return McpProtocolHandler._error_response(request_id, -32000, str(exc))
  109. def create_http_handler(gateway_app, rate_limiter=None):
  110. rpc_handler = PublicMcpHttpHandler(gateway_app, rate_limiter=rate_limiter)
  111. class Handler(BaseHTTPRequestHandler):
  112. def log_message(self, format, *args):
  113. # Override to use Python logging instead of stderr
  114. logger.info(f"{self.address_string()} - {format % args}")
  115. def do_GET(self):
  116. if self.path == '/health':
  117. self._write_json({'ok': True})
  118. return
  119. self.send_response(404)
  120. self.end_headers()
  121. def do_POST(self):
  122. client_ip = extract_client_ip(dict(self.headers.items()), self.client_address)
  123. length = int(self.headers.get('Content-Length') or '0')
  124. if self.path != '/mcp':
  125. logger.warning(f"[HTTP] 404: ip={client_ip}, path={self.path}")
  126. if length > 0:
  127. self.rfile.read(length)
  128. self.send_response(404)
  129. self.end_headers()
  130. return
  131. body = self.rfile.read(length).decode('utf-8-sig')
  132. try:
  133. message = json.loads(body)
  134. except json.JSONDecodeError as exc:
  135. logger.warning(f"[HTTP] invalid JSON: ip={client_ip}, error={exc}")
  136. self._write_json(McpProtocolHandler._error_response(None, -32700, 'Parse error'))
  137. return
  138. method = message.get('method', '')
  139. request_id = message.get('id') if message.get('id') is not None else ''
  140. logger.info(f"[HTTP] request: ip={client_ip}, method={method}, id={request_id}")
  141. response = rpc_handler.handle_json_rpc(dict(self.headers.items()), message, client_ip)
  142. self._write_json(response)
  143. def _write_json(self, payload):
  144. raw = json.dumps(payload, ensure_ascii=False).encode('utf-8')
  145. self.send_response(200)
  146. self.send_header('Content-Type', 'application/json; charset=utf-8')
  147. self.send_header('Content-Length', str(len(raw)))
  148. self.end_headers()
  149. self.wfile.write(raw)
  150. return Handler
  151. def serve_public(gateway_app, host='0.0.0.0', port=8765, enable_rate_limit=True,
  152. rate_limit_max_requests=60, rate_limit_window_seconds=60):
  153. # Configure logging
  154. logging.basicConfig(
  155. level=logging.INFO,
  156. format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
  157. datefmt='%Y-%m-%d %H:%M:%S'
  158. )
  159. # Configure rate limiting
  160. rate_limiter = None
  161. if enable_rate_limit:
  162. rate_limiter = SimpleRateLimiter(
  163. max_requests=rate_limit_max_requests,
  164. window_seconds=rate_limit_window_seconds,
  165. )
  166. logger.info(f"Rate limiting enabled: {rate_limit_max_requests} requests/{rate_limit_window_seconds}s per session (tools/call) / per IP (tools/list)")
  167. logger.info(f"Starting public MCP Gateway on {host}:{port}")
  168. server = ThreadingHTTPServer((host, int(port)), create_http_handler(gateway_app, rate_limiter))
  169. # Schedule periodic cleanup to prevent unbounded memory growth in the rate limiter
  170. if rate_limiter is not None:
  171. def _cleanup_loop():
  172. while True:
  173. time.sleep(300)
  174. # Use the actual window as max_age to avoid deleting entries still within the window
  175. rate_limiter.cleanup(max_age_seconds=rate_limit_window_seconds)
  176. t = threading.Thread(target=_cleanup_loop, daemon=True)
  177. t.start()
  178. try:
  179. server.serve_forever()
  180. except KeyboardInterrupt:
  181. logger.info("Shutting down public MCP Gateway")
  182. server.shutdown()