public_server.py 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146
  1. import json
  2. import logging
  3. import uuid
  4. from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
  5. from constants import DEVICE_INVALID_MESSAGE
  6. from mcp_protocol import McpProtocolHandler
  7. from services.request_context import RequestContextParser
  8. from utils.rate_limiter import SimpleRateLimiter
  9. logger = logging.getLogger(__name__)
  10. def extract_client_ip(headers, client_address):
  11. if client_address and len(client_address) > 0:
  12. return str(client_address[0])
  13. return ''
  14. class PublicMcpHttpHandler:
  15. def __init__(self, gateway_app, context_parser=None, rate_limiter=None):
  16. self.gateway_app = gateway_app
  17. self.context_parser = context_parser or RequestContextParser()
  18. self.rate_limiter = rate_limiter
  19. def handle_json_rpc(self, headers, message, client_ip=''):
  20. request_id = message.get('id') if isinstance(message, dict) else None
  21. method = str((message or {}).get('method') or '').strip()
  22. # Rate limiting by IP
  23. if self.rate_limiter and client_ip:
  24. if not self.rate_limiter.is_allowed(client_ip):
  25. logger.warning(f"[RATE_LIMIT] IP rate limit exceeded: ip={client_ip}, method={method}")
  26. return McpProtocolHandler._error_response(request_id, -32000, 'Rate limit exceeded. Please try again later.')
  27. try:
  28. if method == 'initialize':
  29. return McpProtocolHandler._success_response(request_id, {
  30. 'protocolVersion': McpProtocolHandler.protocol_version,
  31. 'capabilities': {'tools': {'listChanged': False}},
  32. 'serverInfo': {
  33. 'name': McpProtocolHandler.server_name,
  34. 'version': McpProtocolHandler.server_version,
  35. },
  36. })
  37. if method == 'tools/list':
  38. tools = []
  39. for tool in self.gateway_app.list_tools():
  40. normalized = dict(tool)
  41. if 'input_schema' in normalized:
  42. normalized['inputSchema'] = normalized.pop('input_schema')
  43. tools.append(normalized)
  44. return McpProtocolHandler._success_response(request_id, {'tools': tools})
  45. if method == 'tools/call':
  46. context = self.context_parser.parse(headers or {})
  47. if not context.has_session():
  48. raise RuntimeError(DEVICE_INVALID_MESSAGE)
  49. params = message.get('params') or {}
  50. result = self.gateway_app.call_tool(
  51. context.gateway_session_id,
  52. params.get('name'),
  53. params.get('arguments') or {},
  54. request_id='rq_http_{0}_{1}'.format(request_id, uuid.uuid4().hex[:16]),
  55. )
  56. structured_content = result.get('data') or {}
  57. return McpProtocolHandler._success_response(request_id, {
  58. 'content': [{'type': 'text', 'text': McpProtocolHandler._render_text(structured_content)}],
  59. 'structuredContent': structured_content,
  60. 'isError': False,
  61. })
  62. return McpProtocolHandler._error_response(request_id, -32601, 'Method not found: {0}'.format(method))
  63. except Exception as exc:
  64. if method == 'tools/call':
  65. return McpProtocolHandler._success_response(request_id, {
  66. 'content': [{'type': 'text', 'text': str(exc)}],
  67. 'isError': True,
  68. })
  69. return McpProtocolHandler._error_response(request_id, -32000, str(exc))
  70. def create_http_handler(gateway_app, rate_limiter=None):
  71. rpc_handler = PublicMcpHttpHandler(gateway_app, rate_limiter=rate_limiter)
  72. class Handler(BaseHTTPRequestHandler):
  73. def log_message(self, format, *args):
  74. # Override to use Python logging instead of stderr
  75. logger.info(f"{self.address_string()} - {format % args}")
  76. def do_GET(self):
  77. if self.path == '/health':
  78. self._write_json({'ok': True})
  79. return
  80. self.send_response(404)
  81. self.end_headers()
  82. def do_POST(self):
  83. client_ip = extract_client_ip(dict(self.headers.items()), self.client_address)
  84. if self.path != '/mcp':
  85. logger.warning(f"[HTTP] 404: ip={client_ip}, path={self.path}")
  86. self.send_response(404)
  87. self.end_headers()
  88. return
  89. length = int(self.headers.get('Content-Length') or '0')
  90. body = self.rfile.read(length).decode('utf-8-sig')
  91. message = json.loads(body)
  92. method = message.get('method', '')
  93. request_id = message.get('id', '')
  94. logger.info(f"[HTTP] request: ip={client_ip}, method={method}, id={request_id}")
  95. response = rpc_handler.handle_json_rpc(dict(self.headers.items()), message, client_ip)
  96. self._write_json(response)
  97. def _write_json(self, payload):
  98. raw = json.dumps(payload, ensure_ascii=False).encode('utf-8')
  99. self.send_response(200)
  100. self.send_header('Content-Type', 'application/json; charset=utf-8')
  101. self.send_header('Content-Length', str(len(raw)))
  102. self.end_headers()
  103. self.wfile.write(raw)
  104. return Handler
  105. def serve_public(gateway_app, host='0.0.0.0', port=8765, enable_rate_limit=True):
  106. # Configure logging
  107. logging.basicConfig(
  108. level=logging.INFO,
  109. format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
  110. datefmt='%Y-%m-%d %H:%M:%S'
  111. )
  112. # Configure rate limiting (default: 60 requests per minute per IP)
  113. rate_limiter = None
  114. if enable_rate_limit:
  115. rate_limiter = SimpleRateLimiter(max_requests=60, window_seconds=60)
  116. logger.info("Rate limiting enabled: 60 requests/minute per IP")
  117. logger.info(f"Starting public MCP Gateway on {host}:{port}")
  118. server = ThreadingHTTPServer((host, int(port)), create_http_handler(gateway_app, rate_limiter))
  119. try:
  120. server.serve_forever()
  121. except KeyboardInterrupt:
  122. logger.info("Shutting down public MCP Gateway")
  123. server.shutdown()