import json import logging import threading import time import uuid from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from constants import DEVICE_INVALID_MESSAGE from mcp_protocol import McpProtocolHandler from services.request_context import RequestContextParser from utils.rate_limiter import SimpleRateLimiter logger = logging.getLogger(__name__) def extract_client_ip(headers, client_address): if client_address and len(client_address) > 0: return str(client_address[0]) return '' class PublicMcpHttpHandler: def __init__(self, gateway_app, context_parser=None, rate_limiter=None): self.gateway_app = gateway_app self.context_parser = context_parser or RequestContextParser() self.rate_limiter = rate_limiter # Cache registered tool names at startup for rate-key validation; # unknown names fall back to the bare IP bucket, preventing bucket explosion self._known_tools = frozenset(t.get('name', '') for t in gateway_app.list_tools()) def _check_rate_limit(self, rate_key, method, request_id=None, *, log_identity=None): """Returns an error response if rate limit exceeded, else None. rate_key — the bucket key used by the limiter (session-based for tools/call, IP-based for tools/list) log_identity — optional string shown in warning logs (e.g. client_ip) """ if self.rate_limiter and rate_key and not self.rate_limiter.is_allowed(rate_key): logger.warning( f"[RATE_LIMIT] rate limit exceeded: identity={log_identity or rate_key}, method={method}" ) return McpProtocolHandler._error_response(request_id, -32000, 'Rate limit exceeded. Please try again later.') return None def handle_json_rpc(self, headers, message, client_ip=''): request_id = message.get('id') if isinstance(message, dict) else None method = str((message or {}).get('method') or '').strip() try: if method == 'initialize': # initialize is a stateless handshake that only returns server metadata; # rate-limiting it would block clients from connecting at all, so we skip it. return McpProtocolHandler._success_response(request_id, { 'protocolVersion': McpProtocolHandler.protocol_version, 'capabilities': {'tools': {'listChanged': False}}, 'serverInfo': { 'name': McpProtocolHandler.server_name, 'version': McpProtocolHandler.server_version, }, }) if method == 'tools/list': # tools/list has no session context — use a dedicated IP-based bucket # so it doesn't compete with the per-session tools/call quota blocked = self._check_rate_limit( 'list:{0}'.format(client_ip), method, request_id, log_identity=client_ip, ) if blocked: return blocked tools = [] for tool in self.gateway_app.list_tools(): normalized = dict(tool) if 'input_schema' in normalized: normalized['inputSchema'] = normalized.pop('input_schema') tools.append(normalized) return McpProtocolHandler._success_response(request_id, {'tools': tools}) if method == 'tools/call': # Parse context first — tools/call always requires a valid session context = self.context_parser.parse(headers or {}) if not context.has_session(): raise RuntimeError(DEVICE_INVALID_MESSAGE) # Session-based rate limiting: each employee gets an independent quota per tool. # Unknown tool names fall back to the bare session bucket to prevent key explosion. session_id = context.gateway_session_id tool_name = ((message.get('params') or {}).get('name') or '').strip() rate_key = '{0}:{1}'.format(session_id, tool_name) if tool_name in self._known_tools else session_id blocked = self._check_rate_limit(rate_key, method, request_id, log_identity=client_ip) if blocked: return blocked params = message.get('params') or {} result = self.gateway_app.call_tool( session_id, params.get('name'), params.get('arguments') or {}, request_id='rq_http_{0}_{1}'.format(request_id, uuid.uuid4().hex[:16]), client_ip=client_ip, ) structured_content = result.get('data') or {} return McpProtocolHandler._success_response(request_id, { 'content': [{'type': 'text', 'text': McpProtocolHandler._render_text(structured_content)}], 'structuredContent': structured_content, 'isError': False, }) return McpProtocolHandler._error_response(request_id, -32601, 'Method not found: {0}'.format(method)) except Exception as exc: if method == 'tools/call': return McpProtocolHandler._success_response(request_id, { 'content': [{'type': 'text', 'text': str(exc)}], 'isError': True, }) return McpProtocolHandler._error_response(request_id, -32000, str(exc)) def create_http_handler(gateway_app, rate_limiter=None): rpc_handler = PublicMcpHttpHandler(gateway_app, rate_limiter=rate_limiter) class Handler(BaseHTTPRequestHandler): def log_message(self, format, *args): # Override to use Python logging instead of stderr logger.info(f"{self.address_string()} - {format % args}") def do_GET(self): if self.path == '/health': self._write_json({'ok': True}) return self.send_response(404) self.end_headers() def do_POST(self): client_ip = extract_client_ip(dict(self.headers.items()), self.client_address) if self.path != '/mcp': logger.warning(f"[HTTP] 404: ip={client_ip}, path={self.path}") self.send_response(404) self.end_headers() return length = int(self.headers.get('Content-Length') or '0') body = self.rfile.read(length).decode('utf-8-sig') try: message = json.loads(body) except json.JSONDecodeError as exc: logger.warning(f"[HTTP] invalid JSON: ip={client_ip}, error={exc}") self._write_json(McpProtocolHandler._error_response(None, -32700, 'Parse error')) return method = message.get('method', '') request_id = message.get('id') if message.get('id') is not None else '' logger.info(f"[HTTP] request: ip={client_ip}, method={method}, id={request_id}") response = rpc_handler.handle_json_rpc(dict(self.headers.items()), message, client_ip) self._write_json(response) def _write_json(self, payload): raw = json.dumps(payload, ensure_ascii=False).encode('utf-8') self.send_response(200) self.send_header('Content-Type', 'application/json; charset=utf-8') self.send_header('Content-Length', str(len(raw))) self.end_headers() self.wfile.write(raw) return Handler def serve_public(gateway_app, host='0.0.0.0', port=8765, enable_rate_limit=True, rate_limit_max_requests=60, rate_limit_window_seconds=60): # Configure logging logging.basicConfig( level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s', datefmt='%Y-%m-%d %H:%M:%S' ) # Configure rate limiting rate_limiter = None if enable_rate_limit: rate_limiter = SimpleRateLimiter( max_requests=rate_limit_max_requests, window_seconds=rate_limit_window_seconds, ) logger.info(f"Rate limiting enabled: {rate_limit_max_requests} requests/{rate_limit_window_seconds}s per session (tools/call) / per IP (tools/list)") logger.info(f"Starting public MCP Gateway on {host}:{port}") server = ThreadingHTTPServer((host, int(port)), create_http_handler(gateway_app, rate_limiter)) # Schedule periodic cleanup to prevent unbounded memory growth in the rate limiter if rate_limiter is not None: def _cleanup_loop(): while True: time.sleep(300) # Use the actual window as max_age to avoid deleting entries still within the window rate_limiter.cleanup(max_age_seconds=rate_limit_window_seconds) t = threading.Thread(target=_cleanup_loop, daemon=True) t.start() try: server.serve_forever() except KeyboardInterrupt: logger.info("Shutting down public MCP Gateway") server.shutdown()