Parcourir la source

mcp 限流优化

jackson il y a 2 semaines
Parent
commit
96f864bcf3

+ 8 - 1
app.py

@@ -145,7 +145,14 @@ class GatewayApp:
                 session_store=session_store,
                 api_client=ScopedApiClient(config.tools_base_url, timeout=config.timeout_seconds),
             )
-            return serve_public(public_app, host=args.host, port=args.port)
+            return serve_public(
+                public_app,
+                host=args.host,
+                port=args.port,
+                enable_rate_limit=config.rate_limit_enabled,
+                rate_limit_max_requests=config.rate_limit_max_requests,
+                rate_limit_window_seconds=config.rate_limit_window_seconds,
+            )
         elif args.command == 'call':
             tool_args = {
                 'page': args.page,

+ 39 - 1
config.py

@@ -20,6 +20,9 @@ class GatewayConfig:
     session_key: str = ''
     gateway_mode: str = 'local'
     gateway_session_ttl_seconds: int = 2592000
+    rate_limit_enabled: bool = True
+    rate_limit_max_requests: int = 60
+    rate_limit_window_seconds: int = 60
 
     @classmethod
     def from_env(cls, env=None, dotenv_path=''):
@@ -67,6 +70,9 @@ class GatewayConfig:
             session_key=session_key or cls._build_default_session_key(primary_env),
             gateway_mode=gateway_mode,
             gateway_session_ttl_seconds=int(cls._pick(dotenv_env, primary_env, 'FMS_GATEWAY_SESSION_TTL_SECONDS', 'MCP_GATEWAY_SESSION_TTL_SECONDS') or '2592000'),
+            rate_limit_enabled=cls._parse_bool(cls._pick(dotenv_env, primary_env, 'FMS_RATE_LIMIT_ENABLED', 'MCP_RATE_LIMIT_ENABLED'), default=True),
+            rate_limit_max_requests=cls._parse_int(cls._pick(dotenv_env, primary_env, 'FMS_RATE_LIMIT_MAX_REQUESTS', 'MCP_RATE_LIMIT_MAX_REQUESTS'), default=60),
+            rate_limit_window_seconds=cls._parse_int(cls._pick(dotenv_env, primary_env, 'FMS_RATE_LIMIT_WINDOW_SECONDS', 'MCP_RATE_LIMIT_WINDOW_SECONDS'), default=60),
         )
 
     @staticmethod
@@ -111,6 +117,31 @@ class GatewayConfig:
                 return value
         return ''
 
+    @staticmethod
+    def _strip_comment(value):
+        """Strip inline comment from a raw env value (space+# pattern)."""
+        pos = str(value or '').find(' #')
+        return value[:pos].strip() if pos >= 0 else value
+
+    @staticmethod
+    def _parse_int(raw, default):
+        """Parse int from env value, tolerating inline comments from OS env vars."""
+        value = GatewayConfig._strip_comment(str(raw or '').strip())
+        if not value:
+            return default
+        try:
+            return int(value)
+        except ValueError:
+            return default
+
+    @staticmethod
+    def _parse_bool(raw, default=True):
+        """Parse bool from env value; empty / unset → default."""
+        value = GatewayConfig._strip_comment(str(raw or '').strip()).lower()
+        if not value:
+            return default
+        return value not in ('0', 'false', 'no', 'off')
+
     @staticmethod
     def _build_default_session_key(env):
         computer = env.get('COMPUTERNAME') or env.get('HOSTNAME') or os.environ.get('COMPUTERNAME') or os.environ.get('HOSTNAME') or 'unknown-computer'
@@ -131,7 +162,14 @@ class GatewayConfig:
                     continue
                 key, value = line.split('=', 1)
                 key = key.strip()
-                value = value.strip().strip('"').strip("'")
+                value = value.strip()
+                # Strip inline comments for unquoted values (e.g. KEY=123  # comment)
+                if value and value[0] not in ('"', "'"):
+                    comment_pos = value.find(' #')
+                    if comment_pos >= 0:
+                        value = value[:comment_pos].strip()
+                else:
+                    value = value.strip('"').strip("'")
                 if key:
                     data[key] = value
         return data

+ 2 - 1
public_gateway.py

@@ -26,7 +26,7 @@ class PublicGatewayApp:
         request_id = str(request_id or '').strip()
         return request_id or 'rq_{0}'.format(uuid.uuid4().hex[:16])
 
-    def call_tool(self, gateway_session_id, name, arguments=None, request_id=''):
+    def call_tool(self, gateway_session_id, name, arguments=None, request_id='', client_ip=''):
         if name not in self._tools:
             raise KeyError('tool not registered: {0}'.format(name))
 
@@ -49,6 +49,7 @@ class PublicGatewayApp:
                 route_path=tool.route_path,
                 payload=arguments or {},
                 request_id=request_id,
+                client_ip=client_ip,
             )
             if hasattr(self.session_store, 'touch_session'):
                 self.session_store.touch_session(gateway_session_id)

+ 52 - 12
public_server.py

@@ -1,5 +1,7 @@
 import json
 import logging
+import threading
+import time
 import uuid
 from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
 
@@ -22,19 +24,25 @@ class PublicMcpHttpHandler:
         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, client_ip, rate_key, method, request_id=None):
+        """Returns an error response if rate limit exceeded, else None."""
+        if self.rate_limiter and client_ip and not self.rate_limiter.is_allowed(rate_key):
+            logger.warning(f"[RATE_LIMIT] rate limit exceeded: ip={client_ip}, 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()
 
-        # Rate limiting by IP
-        if self.rate_limiter and client_ip:
-            if not self.rate_limiter.is_allowed(client_ip):
-                logger.warning(f"[RATE_LIMIT] IP rate limit exceeded: ip={client_ip}, method={method}")
-                return McpProtocolHandler._error_response(request_id, -32000, 'Rate limit exceeded. Please try again later.')
-
         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}},
@@ -44,6 +52,10 @@ class PublicMcpHttpHandler:
                     },
                 })
             if method == 'tools/list':
+                # Use a dedicated bucket for list so it doesn't compete with tools/call quota
+                blocked = self._check_rate_limit(client_ip, 'list:{0}'.format(client_ip), method, request_id)
+                if blocked:
+                    return blocked
                 tools = []
                 for tool in self.gateway_app.list_tools():
                     normalized = dict(tool)
@@ -52,6 +64,13 @@ class PublicMcpHttpHandler:
                     tools.append(normalized)
                 return McpProtocolHandler._success_response(request_id, {'tools': tools})
             if method == 'tools/call':
+                # Per-tool independent quota: only registered tool names get their own bucket;
+                # unrecognized names fall into the shared IP bucket to prevent bucket explosion
+                tool_name = ((message.get('params') or {}).get('name') or '').strip()
+                rate_key = '{0}:{1}'.format(client_ip, tool_name) if tool_name in self._known_tools else client_ip
+                blocked = self._check_rate_limit(client_ip, rate_key, method, request_id)
+                if blocked:
+                    return blocked
                 context = self.context_parser.parse(headers or {})
                 if not context.has_session():
                     raise RuntimeError(DEVICE_INVALID_MESSAGE)
@@ -61,6 +80,7 @@ class PublicMcpHttpHandler:
                     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, {
@@ -102,9 +122,14 @@ def create_http_handler(gateway_app, rate_limiter=None):
                 return
             length = int(self.headers.get('Content-Length') or '0')
             body = self.rfile.read(length).decode('utf-8-sig')
-            message = json.loads(body)
+            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', '')
+            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)
@@ -120,7 +145,8 @@ def create_http_handler(gateway_app, rate_limiter=None):
     return Handler
 
 
-def serve_public(gateway_app, host='0.0.0.0', port=8765, enable_rate_limit=True):
+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,
@@ -128,14 +154,28 @@ def serve_public(gateway_app, host='0.0.0.0', port=8765, enable_rate_limit=True)
         datefmt='%Y-%m-%d %H:%M:%S'
     )
 
-    # Configure rate limiting (default: 60 requests per minute per IP)
+    # Configure rate limiting
     rate_limiter = None
     if enable_rate_limit:
-        rate_limiter = SimpleRateLimiter(max_requests=60, window_seconds=60)
-        logger.info("Rate limiting enabled: 60 requests/minute per IP")
+        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 IP")
 
     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:

+ 5 - 2
services/scoped_api_client.py

@@ -1,4 +1,4 @@
-from services.api_client import JsonTransport
+from services.api_client import JsonTransport
 
 
 class ScopedApiClient:
@@ -7,7 +7,7 @@ class ScopedApiClient:
         self.transport = transport or JsonTransport()
         self.timeout = int(timeout)
 
-    def call_tool(self, token, tool_code, route_path, payload, request_id):
+    def call_tool(self, token, tool_code, route_path, payload, request_id, client_ip=''):
         token = str(token or '').strip()
         if not token:
             raise RuntimeError('mcp token missing')
@@ -17,4 +17,7 @@ class ScopedApiClient:
             'X-MCP-Tool-Code': tool_code,
             'X-Request-Id': request_id,
         }
+        client_ip = str(client_ip or '').strip()
+        if client_ip:
+            headers['X-MCP-Client-IP'] = client_ip
         return self.transport.post_json(url, payload, headers, self.timeout)

+ 222 - 0
tests/test_app_coverage.py

@@ -0,0 +1,222 @@
+"""
+Coverage补充测试:app.py
+
+目标路径:
+- GatewayApp.from_config — token_store_type 不支持 → ValueError
+- GatewayApp.run_cli call — else 分支(tool 既非 query_order 也非 query_track,有 keyword)
+- GatewayApp.run_cli call — else 分支(tool 既非 query_order 也非 query_track,无 keyword)
+- main() — 正常调用 list-tools
+"""
+import io
+import json
+import os
+import tempfile
+import unittest
+from unittest.mock import MagicMock, patch
+
+from app import GatewayApp, main
+from services.token_store import FileTokenStore
+from tools.query_order import QueryOrderTool
+
+
+# ---------------------------------------------------------------------------
+# GatewayApp.from_config — 不支持的 store type
+# ---------------------------------------------------------------------------
+
+class FromConfigUnsupportedStoreTest(unittest.TestCase):
+    def test_raises_value_error_for_unknown_store_type(self):
+        config = MagicMock()
+        config.token_store_type = 'memcached'
+        with self.assertRaises(ValueError) as ctx:
+            GatewayApp.from_config(config)
+        self.assertIn('unsupported token store type', str(ctx.exception))
+        self.assertIn('memcached', str(ctx.exception))
+
+    def test_from_config_file_store(self):
+        """file store 分支:正常返回 GatewayApp 实例。"""
+        with tempfile.TemporaryDirectory() as tmp_dir:
+            config = MagicMock()
+            config.token_store_type = 'file'
+            config.token_store_path = os.path.join(tmp_dir, 'tok.json')
+            config.refresh_skew_seconds = 120
+            config.auth_base_url = 'http://auth.test'
+            config.client_type = 'workbuddy'
+            config.timeout_seconds = 5
+            config.session_key = 'test_session'
+            config.tools_base_url = 'http://api.test'
+            app = GatewayApp.from_config(config)
+            self.assertIsInstance(app, GatewayApp)
+
+    def test_from_config_redis_store_uses_provided_redis_client(self):
+        """redis store 分支:传入 redis_client 时跳过 RedisSocketClient 创建。"""
+        mock_redis = MagicMock()
+        config = MagicMock()
+        config.token_store_type = 'redis'
+        config.redis_prefix = 'test:'
+        config.session_key = 'sess_1'
+        config.refresh_skew_seconds = 120
+        config.auth_base_url = 'http://auth.test'
+        config.client_type = 'workbuddy'
+        config.timeout_seconds = 5
+        config.tools_base_url = 'http://api.test'
+        # get() 返回 None,避免 RedisTokenStore 在初始化时崩溃
+        mock_redis.get.return_value = None
+        app = GatewayApp.from_config(config, redis_client=mock_redis)
+        self.assertIsInstance(app, GatewayApp)
+
+
+# ---------------------------------------------------------------------------
+# GatewayApp.run_cli call — else 分支
+# ---------------------------------------------------------------------------
+
+class DummyApiClient:
+    def call_tool(self, tool_code, route_path, payload, request_id):
+        return {
+            'code': 'MCP_0000',
+            'data': {},
+            'meta': {'request_id': request_id},
+        }
+
+
+class DummyAuthClient:
+    def __init__(self, token_store):
+        self.token_store = token_store
+
+    def refresh(self, mcp_token):
+        pass
+
+
+def _make_app_with_token():
+    """创建有有效 token 的 GatewayApp,附带一个自定义 fake_tool 工具。"""
+    with tempfile.TemporaryDirectory() as tmp_dir:
+        path = os.path.join(tmp_dir, 'token.json')
+        token_store = FileTokenStore(path)
+        token_store.save('MT_test', '2099-01-01T00:00:00')
+        api_client = DummyApiClient()
+        app = GatewayApp(
+            auth_client=DummyAuthClient(token_store),
+            api_client=api_client,
+            token_store=token_store,
+        )
+        # 注册一个假工具,使 else 分支可以正常 call_tool
+        app._tools['fake_tool'] = QueryOrderTool(api_client=api_client)
+        return app, tmp_dir
+
+
+class RunCliCallElseBranchTest(unittest.TestCase):
+    def test_call_else_branch_with_keyword_passes_keyword_to_tool(self):
+        """
+        --tool fake_tool(非 query_order / query_track)且 --keyword 有值 →
+        走 else 分支,tool_args 中加入 keyword。
+        """
+        with tempfile.TemporaryDirectory() as tmp_dir:
+            path = os.path.join(tmp_dir, 'token.json')
+            token_store = FileTokenStore(path)
+            token_store.save('MT_test', '2099-01-01T00:00:00')
+            api_client = DummyApiClient()
+            app = GatewayApp(
+                auth_client=DummyAuthClient(token_store),
+                api_client=api_client,
+                token_store=token_store,
+            )
+            app._tools['fake_tool'] = QueryOrderTool(api_client=api_client)
+
+            stdout = io.StringIO()
+            code = app.run_cli(
+                ['call', '--tool', 'fake_tool', '--keyword', 'hello'],
+                stdout=stdout,
+            )
+            self.assertEqual(0, code)
+            payload = json.loads(stdout.getvalue())
+            self.assertEqual('MCP_0000', payload['code'])
+
+    def test_call_else_branch_without_keyword_still_calls_tool(self):
+        """
+        --tool fake_tool(非 query_order / query_track)且 --keyword 为空 →
+        走 else 分支,tool_args 中不含 keyword(由工具自己决定是否报错)。
+        注:fake_tool 是 QueryOrderTool,无 keyword 会抛 ValueError →
+        GatewayApp.run_cli 不 catch,直接抛出。
+        """
+        with tempfile.TemporaryDirectory() as tmp_dir:
+            path = os.path.join(tmp_dir, 'token.json')
+            token_store = FileTokenStore(path)
+            token_store.save('MT_test', '2099-01-01T00:00:00')
+            api_client = DummyApiClient()
+            app = GatewayApp(
+                auth_client=DummyAuthClient(token_store),
+                api_client=api_client,
+                token_store=token_store,
+            )
+
+            # 注册一个不需要 keyword 的假工具
+            class NoKeywordTool:
+                name = 'no_kw_tool'
+                route_path = '/fake'
+                requires_session = False
+
+                def metadata(self):
+                    return {'name': self.name, 'description': '', 'input_schema': {'type': 'object', 'properties': {}}}
+
+                def call(self, request_id='', **_kwargs):
+                    return {'code': 'MCP_0000', 'data': {}}
+
+            app._tools['no_kw_tool'] = NoKeywordTool()
+
+            stdout = io.StringIO()
+            # --keyword 不传,else 分支 keyword 为空 → 不加入 tool_args
+            code = app.run_cli(
+                ['call', '--tool', 'no_kw_tool'],
+                stdout=stdout,
+            )
+            self.assertEqual(0, code)
+
+
+# ---------------------------------------------------------------------------
+# main() — 正常调用
+# ---------------------------------------------------------------------------
+
+class MainFunctionTest(unittest.TestCase):
+    def test_main_list_tools_returns_zero(self):
+        """
+        main() 读取 env、构建 app、调用 run_cli(['list-tools'])。
+        直接 mock GatewayConfig.from_env 返回 file store 配置,避免依赖 .env 文件。
+        """
+        import config as config_module
+        with tempfile.TemporaryDirectory() as tmp_dir:
+            token_path = os.path.join(tmp_dir, 'token.json')
+            fake_config = config_module.GatewayConfig(
+                auth_base_url='http://auth.example.com',
+                tools_base_url='http://api.example.com',
+                token_store_type='file',
+                token_store_path=token_path,
+                session_key='test_machine:test_user',
+            )
+            stdout = io.StringIO()
+            with patch.object(config_module.GatewayConfig, 'from_env', return_value=fake_config):
+                with patch('sys.stdout', stdout):
+                    code = main(['list-tools'])
+
+            self.assertEqual(0, code)
+            tools = json.loads(stdout.getvalue())
+            names = [t['name'] for t in tools]
+            self.assertIn('query_order', names)
+            self.assertIn('query_track', names)
+
+    def test_main_propagates_value_error_for_bad_store(self):
+        """
+        当 token_store_type 为不支持的值时,main() 应抛出 ValueError。
+        """
+        import config as config_module
+        bad_config = config_module.GatewayConfig(
+            auth_base_url='http://auth.example.com',
+            tools_base_url='http://api.example.com',
+            token_store_type='bad_store_type',
+            session_key='host:user',
+        )
+        with self.assertRaises(ValueError):
+            with patch.object(config_module.GatewayConfig, 'from_env', return_value=bad_config):
+                main(['list-tools'])
+
+
+if __name__ == '__main__':
+    unittest.main()

+ 81 - 0
tests/test_config_compat.py

@@ -110,5 +110,86 @@ class GatewayConfigCompatTest(unittest.TestCase):
 
         self.assertEqual(1, config.timeout_seconds)
 
+class RateLimitConfigTest(unittest.TestCase):
+    """Tests for _parse_int / _parse_bool helpers and rate limit config parsing."""
+
+    def _config_from_env(self, env):
+        return GatewayConfig.from_env(
+            env=dict({'FMS_API_BASE': 'http://x.test'}, **env),
+            dotenv_path='missing.env',
+        )
+
+    # --- _parse_int via OS env var ---
+
+    def test_rate_limit_max_requests_reads_from_env(self):
+        config = self._config_from_env({'FMS_RATE_LIMIT_MAX_REQUESTS': '120'})
+        self.assertEqual(120, config.rate_limit_max_requests)
+
+    def test_rate_limit_window_seconds_reads_from_env(self):
+        config = self._config_from_env({'FMS_RATE_LIMIT_WINDOW_SECONDS': '30'})
+        self.assertEqual(30, config.rate_limit_window_seconds)
+
+    def test_rate_limit_max_requests_os_env_with_inline_comment_does_not_crash(self):
+        # OS env vars are not processed by _load_dotenv, so inline comments must be
+        # stripped by _parse_int before int() conversion
+        config = self._config_from_env({'FMS_RATE_LIMIT_MAX_REQUESTS': '60 # max per window'})
+        self.assertEqual(60, config.rate_limit_max_requests)
+
+    def test_rate_limit_window_seconds_invalid_value_falls_back_to_default(self):
+        config = self._config_from_env({'FMS_RATE_LIMIT_WINDOW_SECONDS': 'not_a_number'})
+        self.assertEqual(60, config.rate_limit_window_seconds)
+
+    def test_rate_limit_max_requests_empty_falls_back_to_default(self):
+        config = self._config_from_env({})
+        self.assertEqual(60, config.rate_limit_max_requests)
+
+    # --- _parse_bool via OS env var ---
+
+    def test_rate_limit_enabled_default_is_true(self):
+        config = self._config_from_env({})
+        self.assertTrue(config.rate_limit_enabled)
+
+    def test_rate_limit_enabled_zero_disables(self):
+        config = self._config_from_env({'FMS_RATE_LIMIT_ENABLED': '0'})
+        self.assertFalse(config.rate_limit_enabled)
+
+    def test_rate_limit_enabled_false_string_disables(self):
+        config = self._config_from_env({'FMS_RATE_LIMIT_ENABLED': 'false'})
+        self.assertFalse(config.rate_limit_enabled)
+
+    def test_rate_limit_enabled_off_string_disables(self):
+        config = self._config_from_env({'FMS_RATE_LIMIT_ENABLED': 'off'})
+        self.assertFalse(config.rate_limit_enabled)
+
+    def test_rate_limit_enabled_empty_string_keeps_default_enabled(self):
+        # empty string → _pick returns '' → _parse_bool returns default=True
+        config = self._config_from_env({'FMS_RATE_LIMIT_ENABLED': ''})
+        self.assertTrue(config.rate_limit_enabled)
+
+    def test_rate_limit_enabled_one_enables(self):
+        config = self._config_from_env({'FMS_RATE_LIMIT_ENABLED': '1'})
+        self.assertTrue(config.rate_limit_enabled)
+
+    # --- dotenv inline comment stripping ---
+
+    def test_dotenv_inline_comment_stripped_from_int_value(self):
+        with tempfile.TemporaryDirectory() as tmp:
+            path = os.path.join(tmp, '.env')
+            with open(path, 'w') as f:
+                f.write('FMS_API_BASE=http://x.test\n')
+                f.write('FMS_RATE_LIMIT_MAX_REQUESTS=45 # requests per window\n')
+            config = GatewayConfig.from_env(env={}, dotenv_path=path)
+            self.assertEqual(45, config.rate_limit_max_requests)
+
+    def test_dotenv_inline_comment_stripped_from_bool_value(self):
+        with tempfile.TemporaryDirectory() as tmp:
+            path = os.path.join(tmp, '.env')
+            with open(path, 'w') as f:
+                f.write('FMS_API_BASE=http://x.test\n')
+                f.write('FMS_RATE_LIMIT_ENABLED=0 # disabled for testing\n')
+            config = GatewayConfig.from_env(env={}, dotenv_path=path)
+            self.assertFalse(config.rate_limit_enabled)
+
+
 if __name__ == '__main__':
     unittest.main()

+ 282 - 0
tests/test_mcp_protocol_coverage.py

@@ -0,0 +1,282 @@
+"""
+Coverage补充测试:mcp_protocol.py
+
+目标路径:
+- run_stdio — ValueError 导致的 parse error 响应、空行跳过、非 dict message
+- handle_request — 未知 method、非 tools/call 的异常
+- _render_text — 无 columns/track 字段时的 json.dumps fallback
+- _render_text — records 存在但第一条无 status/content/time 时 fallback
+- _render_track_like_text — location / tracking_number / shipment_id / sub_track 可选字段
+- _render_track_like_text — 空 records + tips
+- _render_table_like_text — 空 records + tips
+"""
+import io
+import json
+import unittest
+
+from mcp_protocol import McpProtocolHandler
+
+
+# ---------------------------------------------------------------------------
+# _render_text 的 fallback 路径
+# ---------------------------------------------------------------------------
+
+class RenderTextFallbackTest(unittest.TestCase):
+    def test_empty_dict_returns_ok(self):
+        self.assertEqual('ok', McpProtocolHandler._render_text({}))
+
+    def test_none_returns_ok(self):
+        self.assertEqual('ok', McpProtocolHandler._render_text(None))
+
+    def test_empty_list_returns_ok(self):
+        # non-dict → 走 "not structured_content" 分支
+        self.assertEqual('ok', McpProtocolHandler._render_text([]))
+
+    def test_dict_without_columns_or_track_fields_falls_back_to_json_dumps(self):
+        """有 records 但第一条没有 status/content/time → 走 json.dumps fallback。"""
+        content = {'records': [{'order_no': 'SO001', 'qty': 10}]}
+        result = McpProtocolHandler._render_text(content)
+        # 应该是 json.dumps 的输出,包含键名
+        self.assertIn('order_no', result)
+        self.assertIn('SO001', result)
+
+    def test_dict_with_no_records_falls_back_to_json_dumps(self):
+        """有 columns 但 records 是 None(不是 list)→ columns 是 list 但 records 不是 → fallback。"""
+        content = {'columns': [{'key': 'no', 'name': '编号'}], 'records': None}
+        result = McpProtocolHandler._render_text(content)
+        self.assertIn('columns', result)
+
+    def test_plain_dict_no_columns_no_records_falls_back_to_json(self):
+        content = {'foo': 'bar', 'count': 42}
+        result = McpProtocolHandler._render_text(content)
+        self.assertIn('foo', result)
+        self.assertIn('bar', result)
+
+    def test_records_empty_list_no_columns_falls_back_to_json(self):
+        """records 是空 list(falsy),columns 是 None → 第二个条件 `records` falsy → fallback。"""
+        content = {'records': []}
+        result = McpProtocolHandler._render_text(content)
+        # json.dumps 输出
+        self.assertIn('records', result)
+
+
+# ---------------------------------------------------------------------------
+# _render_table_like_text 的 tips / 空 records 路径
+# ---------------------------------------------------------------------------
+
+class RenderTableEmptyRecordsTest(unittest.TestCase):
+    def test_empty_records_with_tips_shows_tips(self):
+        content = {
+            'columns': [{'key': 'no', 'name': '编号'}, {'key': 'name', 'name': '名称'}],
+            'records': [],
+            'tips': ['暂无数据', '请调整关键词'],
+        }
+        result = McpProtocolHandler._render_text(content)
+        self.assertIn('表头共 2 列:', result)
+        self.assertIn('1. 编号 (no)', result)
+        self.assertIn('提示:暂无数据;请调整关键词', result)
+
+    def test_records_with_none_value_renders_empty_string(self):
+        content = {
+            'columns': [{'key': 'k', 'name': '键'}],
+            'records': [{'k': None}],
+        }
+        result = McpProtocolHandler._render_text(content)
+        self.assertIn('- 键:', result)
+
+
+# ---------------------------------------------------------------------------
+# _render_track_like_text 的可选字段
+# ---------------------------------------------------------------------------
+
+class RenderTrackOptionalFieldsTest(unittest.TestCase):
+    """每个可选字段(location / tracking_number / shipment_id / sub_track)单独覆盖。"""
+
+    def _call(self, **extra):
+        record = {'status': '清关放行', 'content': '启运港放行', 'time': '2026-07-07 10:00:00'}
+        record.update(extra)
+        content = {'records': [record]}
+        return McpProtocolHandler._render_text(content)
+
+    def test_location_rendered_when_present(self):
+        result = self._call(location='宁波港')
+        self.assertIn('- 地点: 宁波港', result)
+
+    def test_location_omitted_when_empty(self):
+        result = self._call(location='')
+        self.assertNotIn('地点', result)
+
+    def test_location_omitted_when_absent(self):
+        result = self._call()
+        self.assertNotIn('地点', result)
+
+    def test_tracking_number_rendered_when_present(self):
+        result = self._call(tracking_number='TN1234567890')
+        self.assertIn('- 快递单号: TN1234567890', result)
+
+    def test_tracking_number_omitted_when_empty(self):
+        result = self._call(tracking_number='')
+        self.assertNotIn('快递单号', result)
+
+    def test_shipment_id_rendered_when_present(self):
+        result = self._call(shipment_id='SHP-9900')
+        self.assertIn('- Shipment ID: SHP-9900', result)
+
+    def test_shipment_id_omitted_when_empty(self):
+        result = self._call(shipment_id='')
+        self.assertNotIn('Shipment ID', result)
+
+    def test_sub_track_rendered_when_truthy(self):
+        result = self._call(sub_track=1)
+        self.assertIn('- 子单轨迹: 是', result)
+
+    def test_sub_track_omitted_when_zero(self):
+        result = self._call(sub_track=0)
+        self.assertNotIn('子单轨迹', result)
+
+    def test_sub_track_omitted_when_absent(self):
+        result = self._call()
+        self.assertNotIn('子单轨迹', result)
+
+    def test_all_optional_fields_together(self):
+        result = self._call(
+            location='上海港',
+            tracking_number='YT123',
+            shipment_id='SHP001',
+            sub_track=2,
+        )
+        self.assertIn('地点: 上海港', result)
+        self.assertIn('快递单号: YT123', result)
+        self.assertIn('Shipment ID: SHP001', result)
+        self.assertIn('子单轨迹: 是', result)
+
+    def test_track_with_summary_and_tips(self):
+        content = {
+            'summary': '共1条轨迹',
+            'records': [
+                {'status': '已签收', 'content': '已签收', 'time': '2026-07-08 12:00:00'}
+            ],
+            'tips': ['仅展示最新轨迹'],
+        }
+        result = McpProtocolHandler._render_text(content)
+        self.assertIn('共1条轨迹', result)
+        self.assertIn('已签收', result)
+        self.assertIn('提示:仅展示最新轨迹', result)
+
+    def test_multiple_track_records(self):
+        content = {
+            'records': [
+                {'status': '已发货', 'content': '揽件', 'time': '2026-07-06 09:00:00'},
+                {'status': '已签收', 'content': '签收', 'time': '2026-07-08 12:00:00'},
+            ],
+        }
+        result = McpProtocolHandler._render_text(content)
+        self.assertIn('轨迹 1:', result)
+        self.assertIn('轨迹 2:', result)
+        self.assertIn('已发货', result)
+        self.assertIn('已签收', result)
+
+
+# ---------------------------------------------------------------------------
+# run_stdio — parse error 分支
+# ---------------------------------------------------------------------------
+
+class RunStdioParseErrorTest(unittest.TestCase):
+    def test_invalid_json_emits_parse_error_response(self):
+        """无效 JSON 行 → 写出 -32700 Parse error 响应。"""
+        handler = McpProtocolHandler(None)
+        stdin = io.StringIO('not-valid-json\n')
+        stdout = io.StringIO()
+
+        handler.run_stdio(stdin=stdin, stdout=stdout)
+
+        line = stdout.getvalue().strip()
+        self.assertTrue(line, 'expected a response line on parse error')
+        response = json.loads(line)
+        self.assertIn('error', response)
+        self.assertEqual(-32700, response['error']['code'])
+        self.assertEqual('Parse error', response['error']['message'])
+        self.assertIsNone(response['id'])
+
+    def test_blank_lines_are_skipped(self):
+        handler = McpProtocolHandler(None)
+        stdin = io.StringIO('\n   \n\t\n')
+        stdout = io.StringIO()
+
+        handler.run_stdio(stdin=stdin, stdout=stdout)
+
+        self.assertEqual('', stdout.getvalue().strip())
+
+    def test_mixed_valid_and_invalid_lines(self):
+        """有效请求和无效行混合;两者都产生输出。"""
+        handler = McpProtocolHandler(None)
+        stdin = io.StringIO(
+            'bad-json\n'
+            + json.dumps({'jsonrpc': '2.0', 'id': 1, 'method': 'initialize', 'params': {}})
+            + '\n'
+        )
+        stdout = io.StringIO()
+        # initialize 需要 gateway_app,这里简单 patch
+        from unittest.mock import MagicMock
+        handler.gateway_app = MagicMock()
+        handler.gateway_app.list_tools.return_value = []
+
+        handler.run_stdio(stdin=stdin, stdout=stdout)
+
+        lines = [l for l in stdout.getvalue().splitlines() if l.strip()]
+        # parse error + initialize response = 2 lines
+        self.assertEqual(2, len(lines))
+        parse_err = json.loads(lines[0])
+        self.assertEqual(-32700, parse_err['error']['code'])
+
+
+# ---------------------------------------------------------------------------
+# handle_message / handle_request 边界情况
+# ---------------------------------------------------------------------------
+
+class HandleMessageEdgeCasesTest(unittest.TestCase):
+    def test_non_dict_message_returns_invalid_request_error(self):
+        handler = McpProtocolHandler(None)
+        response = handler.handle_message('not a dict')
+        self.assertEqual(-32600, response['error']['code'])
+
+    def test_notification_with_unknown_method_returns_none(self):
+        handler = McpProtocolHandler(None)
+        # 无 id → 通知,非 notifications/initialized → returns None
+        response = handler.handle_message({'jsonrpc': '2.0', 'method': 'some/notification'})
+        self.assertIsNone(response)
+
+    def test_unknown_request_method_returns_method_not_found(self):
+        handler = McpProtocolHandler(None)
+        response = handler.handle_request({'id': 1, 'method': 'unknown/method'})
+        self.assertEqual(-32601, response['error']['code'])
+        self.assertIn('Method not found', response['error']['message'])
+
+    def test_non_tools_call_exception_returns_error(self):
+        """tools/list 抛异常 → 返回 error 响应(非 isError 内容)。"""
+        from unittest.mock import MagicMock
+        handler = McpProtocolHandler(MagicMock())
+        handler.gateway_app.list_tools.side_effect = RuntimeError('backend error')
+
+        response = handler.handle_request({'id': 5, 'method': 'tools/list'})
+        self.assertIn('error', response)
+        self.assertIn('backend error', response['error']['message'])
+
+    def test_tools_call_exception_returns_is_error_content(self):
+        """tools/call 抛异常 → 返回 isError: true 的 content。"""
+        from unittest.mock import MagicMock
+        handler = McpProtocolHandler(MagicMock())
+        handler.gateway_app.call_tool.side_effect = RuntimeError('tool failed')
+
+        response = handler.handle_request({
+            'id': 6,
+            'method': 'tools/call',
+            'params': {'name': 'query_order', 'arguments': {}},
+        })
+        self.assertEqual(False, 'error' in response)
+        self.assertTrue(response['result']['isError'])
+        self.assertIn('tool failed', response['result']['content'][0]['text'])
+
+
+if __name__ == '__main__':
+    unittest.main()

+ 12 - 2
tests/test_public_gateway.py

@@ -19,8 +19,8 @@ class FakeApiClient:
     def __init__(self):
         self.calls = []
 
-    def call_tool(self, token, tool_code, route_path, payload, request_id):
-        self.calls.append((token, tool_code, route_path, payload, request_id))
+    def call_tool(self, token, tool_code, route_path, payload, request_id, client_ip=''):
+        self.calls.append((token, tool_code, route_path, payload, request_id, client_ip))
         return {'code': 'MCP_0000', 'data': {'token_used': token}}
 
 
@@ -57,6 +57,16 @@ class PublicGatewayAppTest(unittest.TestCase):
 
         self.assertIn('这台设备的 Workbuddy 配置已失效,请重新生成配置', str(error.exception))
 
+    def test_tool_call_forwards_client_ip_to_api_client(self):
+        store = FakeSessionStore()
+        store.sessions['GWS_A'] = {'mcp_token': 'MT_A'}
+        api_client = FakeApiClient()
+        app = PublicGatewayApp(session_store=store, api_client=api_client, auth_client=None)
+
+        app.call_tool('GWS_A', 'query_order', {'keyword': 'A'}, request_id='rq_a', client_ip='203.0.113.9')
+
+        self.assertEqual('203.0.113.9', api_client.calls[0][5])
+
     def test_successful_tool_call_touches_gateway_session(self):
         store = FakeSessionStore()
         store.sessions['GWS_A'] = {'mcp_token': 'MT_A'}

+ 14 - 0
tests/test_public_gateway_unit.py

@@ -129,6 +129,20 @@ class TestPublicGatewayApp(unittest.TestCase):
         call_args = self.mock_api_client.call_tool.call_args[1]
         self.assertEqual(call_args['request_id'], custom_request_id)
 
+    def test_call_tool_forwards_client_ip_to_api_client(self):
+        self.mock_session_store.get.return_value = {
+            'mcp_token': 'MT_token',
+            'admin_id': 1,
+            'company_id': 1
+        }
+
+        self.mock_api_client.call_tool.return_value = {'code': '0'}
+
+        self.app.call_tool('GWS_test', 'query_order', {}, 'rq_client_ip', client_ip='203.0.113.9')
+
+        call_args = self.mock_api_client.call_tool.call_args[1]
+        self.assertEqual(call_args['client_ip'], '203.0.113.9')
+
     def test_call_tool_logs_error_on_exception(self):
         gateway_session_id = 'GWS_error_test'
         tool_name = 'query_order'

+ 83 - 2
tests/test_public_server.py

@@ -1,6 +1,7 @@
 import unittest
 
 from public_server import PublicMcpHttpHandler, extract_client_ip
+from utils.rate_limiter import SimpleRateLimiter
 
 
 class FakeContext:
@@ -23,8 +24,8 @@ class FakeGateway:
     def list_tools(self):
         return [{'name': 'query_order', 'description': 'query order', 'input_schema': {'type': 'object'}}]
 
-    def call_tool(self, gateway_session_id, name, arguments=None, request_id=''):
-        self.calls.append((gateway_session_id, name, arguments, request_id))
+    def call_tool(self, gateway_session_id, name, arguments=None, request_id='', client_ip=''):
+        self.calls.append((gateway_session_id, name, arguments, request_id, client_ip))
         return {'code': 'MCP_0000', 'data': {'ok': True}}
 
 
@@ -43,11 +44,13 @@ class PublicMcpHttpHandlerTest(unittest.TestCase):
                     'arguments': {'keyword': 'USC'},
                 },
             },
+            client_ip='10.0.0.5'
         )
 
         self.assertEqual(False, response['result']['isError'])
         self.assertEqual('GWS_A', gateway.calls[0][0])
         self.assertEqual('query_order', gateway.calls[0][1])
+        self.assertEqual('10.0.0.5', gateway.calls[0][4])
 
     def test_missing_session_on_tool_call_returns_error_content(self):
         gateway = FakeGateway()
@@ -73,6 +76,84 @@ class PublicMcpHttpHandlerTest(unittest.TestCase):
 
         self.assertEqual('10.0.0.5', client_ip)
 
+
+class RateLimitTest(unittest.TestCase):
+    def _make_handler(self, max_requests=2):
+        gateway = FakeGateway()
+        limiter = SimpleRateLimiter(max_requests=max_requests, window_seconds=60)
+        return PublicMcpHttpHandler(gateway, context_parser=FakeParser(), rate_limiter=limiter)
+
+    def _tools_call_msg(self, tool_name='query_order'):
+        return {
+            'jsonrpc': '2.0', 'id': 1,
+            'method': 'tools/call',
+            'params': {'name': tool_name, 'arguments': {'keyword': 'test'}},
+        }
+
+    # --- tools/call per-tool独立限流 ---
+
+    def test_known_tool_rate_limited_after_quota_exhausted(self):
+        handler = self._make_handler(max_requests=1)
+        handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4')
+        response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4')
+        self.assertIn('error', response)
+        self.assertIn('Rate limit', response['error']['message'])
+
+    def test_two_known_tools_have_independent_quotas(self):
+        # query_order 限流不影响 query_track
+        handler = self._make_handler(max_requests=1)
+        handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_order'), client_ip='1.2.3.4')
+        # query_order quota exhausted, query_track should still work
+        response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_track'), client_ip='1.2.3.4')
+        self.assertNotIn('error', response)
+
+    def test_unknown_tool_name_uses_shared_ip_bucket_not_new_bucket(self):
+        # 未知工具名应归入 client_ip 桶,不能无限新建桶
+        handler = self._make_handler(max_requests=1)
+        # 先用 client_ip 桶打一次(用未知工具名触发)
+        handler.handle_json_rpc({}, self._tools_call_msg('fake_tool_1'), client_ip='1.2.3.4')
+        # 再用另一个未知工具名,应该命中同一个IP桶,被限流
+        response = handler.handle_json_rpc({}, self._tools_call_msg('fake_tool_2'), client_ip='1.2.3.4')
+        self.assertIn('error', response)
+        self.assertIn('Rate limit', response['error']['message'])
+
+    def test_different_ips_have_independent_quotas(self):
+        handler = self._make_handler(max_requests=1)
+        handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.1.1.1')
+        response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='2.2.2.2')
+        self.assertNotIn('error', response)
+
+    # --- initialize 不受限流,tools/list 有独立 bucket ---
+
+    def test_initialize_is_never_rate_limited(self):
+        # initialize 是握手协议,无论请求多少次都不应被限流
+        handler = self._make_handler(max_requests=1)
+        for i in range(5):
+            response = handler.handle_json_rpc({}, {'jsonrpc': '2.0', 'id': i, 'method': 'initialize'}, client_ip='1.2.3.4')
+            self.assertNotIn('error', response, f'initialize should never be rate limited (attempt {i})')
+            self.assertIn('protocolVersion', response['result'])
+
+    def test_tools_list_is_rate_limited(self):
+        handler = self._make_handler(max_requests=1)
+        handler.handle_json_rpc({}, {'jsonrpc': '2.0', 'id': 1, 'method': 'tools/list'}, client_ip='1.2.3.4')
+        response = handler.handle_json_rpc({}, {'jsonrpc': '2.0', 'id': 2, 'method': 'tools/list'}, client_ip='1.2.3.4')
+        self.assertIn('error', response)
+        self.assertIn('Rate limit', response['error']['message'])
+
+    def test_tools_list_quota_independent_from_tools_call_quota(self):
+        # tools/list 和 tools/call 使用不同 bucket,互不影响
+        handler = self._make_handler(max_requests=1)
+        # 用掉 tools/list 的1个 slot
+        handler.handle_json_rpc({}, {'jsonrpc': '2.0', 'id': 1, 'method': 'tools/list'}, client_ip='1.2.3.4')
+        # tools/call 应仍可正常执行(使用独立的 client_ip:tool_name bucket)
+        response = handler.handle_json_rpc(
+            {'X-Gateway-Session': 'GWS_A'},
+            {'jsonrpc': '2.0', 'id': 2, 'method': 'tools/call',
+             'params': {'name': 'query_order', 'arguments': {'keyword': 'test'}}},
+            client_ip='1.2.3.4',
+        )
+        self.assertNotIn('error', response)
+
 if __name__ == '__main__':
     unittest.main()
 

+ 173 - 0
tests/test_public_server_coverage.py

@@ -0,0 +1,173 @@
+"""
+Coverage补充测试:public_server.py
+
+目标路径:
+- extract_client_ip — 空 client_address
+- do_GET — 非 /health 路径 → 404
+- do_POST — Content-Type 正确但 body 为无效 JSON → -32700 parse error
+- PublicMcpHttpHandler.handle_json_rpc — 未知 method → -32601 error
+"""
+import json
+import socket
+import threading
+import unittest
+from http.server import HTTPServer
+
+from public_server import PublicMcpHttpHandler, create_http_handler, extract_client_ip
+
+
+# ---------------------------------------------------------------------------
+# extract_client_ip 边界情况
+# ---------------------------------------------------------------------------
+
+class ExtractClientIpTest(unittest.TestCase):
+    def test_returns_empty_string_when_client_address_is_none(self):
+        result = extract_client_ip({}, None)
+        self.assertEqual('', result)
+
+    def test_returns_empty_string_when_client_address_is_empty_tuple(self):
+        result = extract_client_ip({}, ())
+        self.assertEqual('', result)
+
+    def test_returns_ip_from_client_address(self):
+        result = extract_client_ip({'X-Forwarded-For': '1.2.3.4'}, ('10.0.0.1', 12345))
+        self.assertEqual('10.0.0.1', result)
+
+
+# ---------------------------------------------------------------------------
+# PublicMcpHttpHandler.handle_json_rpc — 未知 method
+# ---------------------------------------------------------------------------
+
+class FakeGateway:
+    def list_tools(self):
+        return [{'name': 'query_order', 'description': '', 'input_schema': {'type': 'object'}}]
+
+    def call_tool(self, *args, **kwargs):
+        return {'data': {}}
+
+
+class HandleJsonRpcUnknownMethodTest(unittest.TestCase):
+    def test_unknown_method_returns_method_not_found(self):
+        handler = PublicMcpHttpHandler(FakeGateway())
+        response = handler.handle_json_rpc(
+            headers={},
+            message={'jsonrpc': '2.0', 'id': 1, 'method': 'unknown/method'},
+        )
+        self.assertIn('error', response)
+        self.assertEqual(-32601, response['error']['code'])
+
+
+# ---------------------------------------------------------------------------
+# 集成测试:启动真实 HTTP 服务器,覆盖 do_GET 404 和 do_POST JSON 错误
+# ---------------------------------------------------------------------------
+
+def _find_free_port():
+    with socket.socket() as s:
+        s.bind(('127.0.0.1', 0))
+        return s.getsockname()[1]
+
+
+class HttpHandlerIntegrationTest(unittest.TestCase):
+    """
+    启动一个真实的 ThreadingHTTPServer(随机端口),测试
+    HTTP 层面的边界路径。
+    """
+
+    @classmethod
+    def setUpClass(cls):
+        port = _find_free_port()
+        gateway = FakeGateway()
+        handler_class = create_http_handler(gateway, rate_limiter=None)
+        cls.server = HTTPServer(('127.0.0.1', port), handler_class)
+        cls.server_thread = threading.Thread(target=cls.server.serve_forever, daemon=True)
+        cls.server_thread.start()
+        cls.base_url = 'http://127.0.0.1:{0}'.format(port)
+
+    @classmethod
+    def tearDownClass(cls):
+        cls.server.shutdown()
+
+    def _get(self, path):
+        import urllib.request
+        req = urllib.request.Request(cls_url := self.base_url + path, method='GET')
+        try:
+            with urllib.request.urlopen(req, timeout=5) as resp:
+                return resp.status, resp.read()
+        except Exception as e:
+            # urllib raises HTTPError for non-2xx
+            if hasattr(e, 'code'):
+                return e.code, b''
+            raise
+
+    def _post(self, path, body, content_type='application/json'):
+        import urllib.request
+        data = body.encode('utf-8') if isinstance(body, str) else body
+        req = urllib.request.Request(
+            self.base_url + path,
+            data=data,
+            headers={'Content-Type': content_type, 'Content-Length': str(len(data))},
+            method='POST',
+        )
+        try:
+            with urllib.request.urlopen(req, timeout=5) as resp:
+                return resp.status, json.loads(resp.read())
+        except Exception as e:
+            if hasattr(e, 'code') and hasattr(e, 'read'):
+                body_bytes = e.read()
+                return e.code, json.loads(body_bytes) if body_bytes else {}
+            raise
+
+    # do_GET /health → 200
+    def test_get_health_returns_200(self):
+        status, body = self._get('/health')
+        self.assertEqual(200, status)
+
+    # do_GET 非 /health → 404(覆盖 lines 113-114)
+    def test_get_unknown_path_returns_404(self):
+        import urllib.request
+        req = urllib.request.Request(self.base_url + '/not-a-real-path', method='GET')
+        try:
+            urllib.request.urlopen(req, timeout=5)
+            self.fail('Expected 404')
+        except Exception as e:
+            self.assertEqual(404, e.code)
+
+    # do_POST 无效 JSON → -32700 parse error(覆盖 lines 127-130)
+    def test_post_invalid_json_returns_parse_error(self):
+        status, body = self._post('/mcp', 'this is not json at all')
+        self.assertEqual(200, status)
+        self.assertIn('error', body)
+        self.assertEqual(-32700, body['error']['code'])
+        self.assertEqual('Parse error', body['error']['message'])
+
+    # do_POST 到非 /mcp 路径 → 404
+    def test_post_wrong_path_returns_404(self):
+        import urllib.request
+        data = json.dumps({'jsonrpc': '2.0', 'id': 1, 'method': 'initialize'}).encode()
+        req = urllib.request.Request(
+            self.base_url + '/wrong',
+            data=data,
+            headers={'Content-Type': 'application/json', 'Content-Length': str(len(data))},
+            method='POST',
+        )
+        try:
+            urllib.request.urlopen(req, timeout=5)
+            self.fail('Expected 404')
+        except Exception as e:
+            self.assertEqual(404, e.code)
+
+    # do_POST 正常 initialize → 200
+    def test_post_initialize_returns_200(self):
+        status, body = self._post('/mcp', json.dumps({
+            'jsonrpc': '2.0',
+            'id': 1,
+            'method': 'initialize',
+            'params': {},
+        }))
+        self.assertEqual(200, status)
+        self.assertIn('result', body)
+        self.assertIn('protocolVersion', body['result'])
+
+
+if __name__ == '__main__':
+    unittest.main()

+ 87 - 0
tests/test_query_order_coverage.py

@@ -0,0 +1,87 @@
+"""
+Coverage补充测试:tools/query_order.py
+
+目标路径:
+- call() — api_client is None → RuntimeError
+- call() — keyword strip 后为空 → ValueError
+"""
+import unittest
+from unittest.mock import MagicMock
+
+from tools.query_order import QueryOrderTool
+
+
+class QueryOrderToolCoverageTest(unittest.TestCase):
+    # ---- api_client is None ----
+
+    def test_call_raises_when_no_api_client(self):
+        tool = QueryOrderTool(api_client=None)
+        with self.assertRaises(RuntimeError) as ctx:
+            tool.call('some keyword')
+        self.assertIn('api client is required', str(ctx.exception))
+
+    def test_call_raises_when_api_client_not_set(self):
+        tool = QueryOrderTool()
+        with self.assertRaises(RuntimeError):
+            tool.call('kw')
+
+    # ---- keyword 为空 ----
+
+    def test_call_raises_on_empty_string_keyword(self):
+        tool = QueryOrderTool(api_client=MagicMock())
+        with self.assertRaises(ValueError) as ctx:
+            tool.call('')
+        self.assertIn('keyword is required', str(ctx.exception))
+
+    def test_call_raises_on_whitespace_only_keyword(self):
+        tool = QueryOrderTool(api_client=MagicMock())
+        with self.assertRaises(ValueError) as ctx:
+            tool.call('   ')
+        self.assertIn('keyword is required', str(ctx.exception))
+
+    def test_call_raises_on_tab_keyword(self):
+        tool = QueryOrderTool(api_client=MagicMock())
+        with self.assertRaises(ValueError):
+            tool.call('\t')
+
+    # ---- 正常路径(验证 strip 和参数传递)----
+
+    def test_call_strips_keyword_before_passing(self):
+        client = MagicMock()
+        client.call_tool.return_value = {'code': 'MCP_0000', 'data': {}}
+        tool = QueryOrderTool(api_client=client)
+
+        tool.call('  SO001  ', page=2, limit=50, request_id='rq_test')
+
+        _, _, payload, _ = client.call_tool.call_args[0]
+        self.assertEqual('SO001', payload['keyword'])
+
+    def test_call_clamps_limit_to_100(self):
+        client = MagicMock()
+        client.call_tool.return_value = {'code': 'MCP_0000', 'data': {}}
+        tool = QueryOrderTool(api_client=client)
+
+        tool.call('kw', limit=999)
+
+        _, _, payload, _ = client.call_tool.call_args[0]
+        self.assertEqual(100, payload['limit'])
+
+    def test_call_clamps_page_to_1(self):
+        client = MagicMock()
+        client.call_tool.return_value = {'code': 'MCP_0000', 'data': {}}
+        tool = QueryOrderTool(api_client=client)
+
+        tool.call('kw', page=0)
+
+        _, _, payload, _ = client.call_tool.call_args[0]
+        self.assertEqual(1, payload['page'])
+
+    def test_metadata_has_required_keyword(self):
+        tool = QueryOrderTool()
+        meta = tool.metadata()
+        self.assertIn('keyword', meta['input_schema']['required'])
+        self.assertEqual('query_order', meta['name'])
+
+
+if __name__ == '__main__':
+    unittest.main()

+ 15 - 0
tests/test_scoped_api_client.py

@@ -44,6 +44,21 @@ class TestScopedApiClient(unittest.TestCase):
         self.assertEqual(result['code'], '0')
         self.assertEqual(result['data']['order'], 'details')
 
+    def test_call_tool_includes_client_ip_header_when_present(self):
+        self.mock_transport.post_json.return_value = {'code': '0'}
+
+        self.client.call_tool(
+            'MT_valid_token',
+            'query_order',
+            '/mcp/tools/query_order',
+            {'order_no': 'ABC123'},
+            'rq_ip_header',
+            client_ip='203.0.113.9'
+        )
+
+        call_headers = self.mock_transport.post_json.call_args[0][2]
+        self.assertEqual(call_headers['X-MCP-Client-IP'], '203.0.113.9')
+
     def test_call_tool_raises_on_empty_token(self):
         with self.assertRaises(RuntimeError) as context:
             self.client.call_tool('', 'query_order', '/path', {}, 'rq_1')

+ 340 - 0
tests/test_token_store_coverage.py

@@ -0,0 +1,340 @@
+"""
+Coverage补充测试:services/token_store.py
+
+目标路径:
+- InMemoryTokenStore.is_expiring — 无session、无expire_time、Z后缀、过期中、未过期
+- InMemoryTokenStore.require_token — 无session、token为空、正常返回
+- RedisTokenStore — 初始化校验、save(含/不含TTL)、get(bytes/str/空串/None)、clear、is_expiring
+- RedisTokenStore._ttl_seconds — None、Z后缀、naive datetime
+- RedisSocketClient._read_response — +/-/:/$/$(nil)/*(含空数组)/unsupported/empty
+- RedisSocketClient._execute — 带 password+db、不带 password+db
+- RedisSocketClient.set/get/delete(mocked socket)
+"""
+import io
+import json
+import unittest
+from datetime import datetime, timedelta, timezone
+from unittest.mock import MagicMock, patch
+
+from services.token_store import (
+    FileTokenStore,
+    InMemoryTokenStore,
+    RedisSocketClient,
+    RedisTokenStore,
+)
+
+
+# ---------------------------------------------------------------------------
+# InMemoryTokenStore
+# ---------------------------------------------------------------------------
+
+class InMemoryIsExpiringTest(unittest.TestCase):
+    def test_no_session_is_expiring(self):
+        store = InMemoryTokenStore()
+        self.assertTrue(store.is_expiring())
+
+    def test_session_with_none_expire_time_is_expiring(self):
+        store = InMemoryTokenStore()
+        store._session = {'token': 'T', 'expire_time': None}
+        self.assertTrue(store.is_expiring())
+
+    def test_session_with_empty_expire_time_is_expiring(self):
+        store = InMemoryTokenStore()
+        store._session = {'token': 'T', 'expire_time': ''}
+        self.assertTrue(store.is_expiring())
+
+    def test_expire_soon_within_skew_is_expiring(self):
+        store = InMemoryTokenStore(refresh_skew_seconds=120)
+        # 60 秒后到期,在 120s 的 skew 窗口内 → 算作 expiring
+        soon = (datetime.now(timezone.utc) + timedelta(seconds=60)).strftime('%Y-%m-%dT%H:%M:%SZ')
+        store._session = {'token': 'T', 'expire_time': soon}
+        self.assertTrue(store.is_expiring())
+
+    def test_expire_far_future_not_expiring(self):
+        store = InMemoryTokenStore(refresh_skew_seconds=120)
+        store._session = {'token': 'T', 'expire_time': '2099-01-01T00:00:00Z'}
+        self.assertFalse(store.is_expiring())
+
+    def test_expire_time_z_suffix_parsed_correctly(self):
+        store = InMemoryTokenStore(refresh_skew_seconds=0)
+        # 已过期(1秒前)
+        past = (datetime.now(timezone.utc) - timedelta(seconds=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
+        store._session = {'token': 'T', 'expire_time': past}
+        self.assertTrue(store.is_expiring())
+
+    def test_expire_time_without_z_suffix(self):
+        store = InMemoryTokenStore(refresh_skew_seconds=0)
+        # 无时区后缀的 ISO 格式(naive datetime)
+        past = (datetime.now() - timedelta(seconds=1)).strftime('%Y-%m-%dT%H:%M:%S')
+        store._session = {'token': 'T', 'expire_time': past}
+        self.assertTrue(store.is_expiring())
+
+
+class InMemoryRequireTokenTest(unittest.TestCase):
+    def test_raises_when_no_session(self):
+        store = InMemoryTokenStore()
+        with self.assertRaises(RuntimeError) as ctx:
+            store.require_token()
+        self.assertIn('mcp token missing', str(ctx.exception))
+
+    def test_raises_when_token_is_empty_string(self):
+        store = InMemoryTokenStore()
+        store._session = {'token': '', 'expire_time': '2099-01-01T00:00:00'}
+        with self.assertRaises(RuntimeError):
+            store.require_token()
+
+    def test_returns_token_when_present(self):
+        store = InMemoryTokenStore()
+        store.save('MT_abc123', '2099-01-01T00:00:00')
+        self.assertEqual('MT_abc123', store.require_token())
+
+
+# ---------------------------------------------------------------------------
+# RedisTokenStore(使用 MagicMock client 隔离真实 Redis)
+# ---------------------------------------------------------------------------
+
+class RedisTokenStoreMockedTest(unittest.TestCase):
+    def _make_store(self, raw_get=None, prefix='test:', session_key='session_1'):
+        client = MagicMock()
+        client.get.return_value = raw_get
+        store = RedisTokenStore(client, prefix=prefix, session_key=session_key)
+        return store, client
+
+    # 初始化
+    def test_empty_session_key_raises_value_error(self):
+        client = MagicMock()
+        with self.assertRaises(ValueError) as ctx:
+            RedisTokenStore(client, session_key='')
+        self.assertIn('session_key is required', str(ctx.exception))
+
+    def test_whitespace_session_key_raises_value_error(self):
+        client = MagicMock()
+        with self.assertRaises(ValueError):
+            RedisTokenStore(client, session_key='   ')
+
+    def test_key_composed_from_prefix_and_session_key(self):
+        store, _ = self._make_store(prefix='ns:', session_key='abc')
+        self.assertEqual('ns:abc', store.key)
+
+    # save
+    def test_save_calls_client_set_with_positive_ttl(self):
+        store, client = self._make_store()
+        store.save('MT_001', '2099-01-01T00:00:00Z')
+        self.assertTrue(client.set.called)
+        args, kwargs = client.set.call_args
+        self.assertEqual('test:session_1', args[0])
+        parsed = json.loads(args[1])
+        self.assertEqual('MT_001', parsed['token'])
+        self.assertIsNotNone(kwargs.get('ex'))
+        self.assertGreater(kwargs['ex'], 0)
+
+    def test_save_with_none_expire_time_passes_no_ex(self):
+        store, client = self._make_store()
+        store.save('MT_001', None)
+        _, kwargs = client.set.call_args
+        self.assertIsNone(kwargs.get('ex'))
+
+    # get
+    def test_get_returns_none_when_redis_returns_none(self):
+        store, _ = self._make_store(raw_get=None)
+        self.assertIsNone(store.get())
+
+    def test_get_returns_none_on_empty_string(self):
+        store, _ = self._make_store(raw_get='')
+        self.assertIsNone(store.get())
+
+    def test_get_decodes_bytes(self):
+        data = {'token': 'MT_bytes', 'expire_time': '2099-01-01T00:00:00'}
+        store, _ = self._make_store(raw_get=json.dumps(data).encode('utf-8'))
+        result = store.get()
+        self.assertEqual('MT_bytes', result['token'])
+
+    def test_get_handles_string_response(self):
+        data = {'token': 'MT_str', 'expire_time': '2099-01-01T00:00:00'}
+        store, _ = self._make_store(raw_get=json.dumps(data))
+        result = store.get()
+        self.assertEqual('MT_str', result['token'])
+
+    # clear
+    def test_clear_calls_client_delete_with_correct_key(self):
+        store, client = self._make_store()
+        store.clear()
+        client.delete.assert_called_once_with('test:session_1')
+        self.assertIsNone(store._session)
+
+    # is_expiring — refreshes from Redis first
+    def test_is_expiring_reads_from_redis_first(self):
+        client = MagicMock()
+        data = {'token': 'MT_live', 'expire_time': '2099-01-01T00:00:00'}
+        client.get.return_value = json.dumps(data)
+        store = RedisTokenStore(client, prefix='t:', session_key='s')
+        self.assertFalse(store.is_expiring())
+        self.assertTrue(client.get.called)
+
+    def test_is_expiring_true_when_redis_has_no_data(self):
+        store, _ = self._make_store(raw_get=None)
+        self.assertTrue(store.is_expiring())
+
+
+class RedisTtlSecondsTest(unittest.TestCase):
+    def test_none_returns_none(self):
+        self.assertIsNone(RedisTokenStore._ttl_seconds(None))
+
+    def test_empty_string_returns_none(self):
+        self.assertIsNone(RedisTokenStore._ttl_seconds(''))
+
+    def test_z_suffix_parsed_as_utc(self):
+        future = (datetime.now(timezone.utc) + timedelta(hours=2)).strftime('%Y-%m-%dT%H:%M:%SZ')
+        ttl = RedisTokenStore._ttl_seconds(future)
+        self.assertGreater(ttl, 3600)
+
+    def test_naive_datetime_uses_local_time(self):
+        future = (datetime.now() + timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M:%S')
+        ttl = RedisTokenStore._ttl_seconds(future)
+        self.assertIsNotNone(ttl)
+        self.assertGreater(ttl, 0)
+
+    def test_already_past_returns_1_at_minimum(self):
+        # 已过期的时间仍返回 max(1, ...)
+        past = (datetime.now(timezone.utc) - timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
+        ttl = RedisTokenStore._ttl_seconds(past)
+        self.assertEqual(1, ttl)
+
+
+# ---------------------------------------------------------------------------
+# RedisSocketClient._read_response(直接传入 BytesIO 流,不需要真实 socket)
+# ---------------------------------------------------------------------------
+
+class RedisReadResponseTest(unittest.TestCase):
+    """_read_response 是实例方法,通过实例调用。"""
+
+    def _client(self):
+        return RedisSocketClient()
+
+    def _stream(self, data: bytes) -> io.BytesIO:
+        return io.BytesIO(data)
+
+    def test_simple_string_ok(self):
+        result = self._client()._read_response(self._stream(b'+OK\r\n'))
+        self.assertEqual('OK', result)
+
+    def test_simple_string_pong(self):
+        result = self._client()._read_response(self._stream(b'+PONG\r\n'))
+        self.assertEqual('PONG', result)
+
+    def test_error_raises_runtime_error(self):
+        with self.assertRaises(RuntimeError) as ctx:
+            self._client()._read_response(self._stream(b'-ERR unknown command\r\n'))
+        self.assertIn('ERR unknown command', str(ctx.exception))
+
+    def test_integer(self):
+        result = self._client()._read_response(self._stream(b':7\r\n'))
+        self.assertEqual(7, result)
+
+    def test_integer_zero(self):
+        result = self._client()._read_response(self._stream(b':0\r\n'))
+        self.assertEqual(0, result)
+
+    def test_bulk_string(self):
+        result = self._client()._read_response(self._stream(b'$5\r\nhello\r\n'))
+        self.assertEqual('hello', result)
+
+    def test_bulk_string_nil(self):
+        result = self._client()._read_response(self._stream(b'$-1\r\n'))
+        self.assertIsNone(result)
+
+    def test_array_two_elements(self):
+        result = self._client()._read_response(self._stream(b'*2\r\n+alpha\r\n+beta\r\n'))
+        self.assertEqual(['alpha', 'beta'], result)
+
+    def test_array_empty(self):
+        result = self._client()._read_response(self._stream(b'*0\r\n'))
+        self.assertEqual([], result)
+
+    def test_array_with_bulk_strings(self):
+        result = self._client()._read_response(self._stream(b'*2\r\n$3\r\nfoo\r\n$3\r\nbar\r\n'))
+        self.assertEqual(['foo', 'bar'], result)
+
+    def test_unsupported_prefix_raises(self):
+        with self.assertRaises(RuntimeError) as ctx:
+            self._client()._read_response(self._stream(b'?nope\r\n'))
+        self.assertIn('unsupported redis response', str(ctx.exception))
+
+    def test_empty_response_raises(self):
+        with self.assertRaises(RuntimeError) as ctx:
+            self._client()._read_response(self._stream(b''))
+        self.assertIn('empty redis response', str(ctx.exception))
+
+
+# ---------------------------------------------------------------------------
+# RedisSocketClient._execute(用 patch mock socket,覆盖 AUTH/SELECT 分支)
+# ---------------------------------------------------------------------------
+
+class RedisSocketClientExecuteTest(unittest.TestCase):
+    """通过 mock socket 验证 _execute 中 AUTH/SELECT 的条件分支。"""
+
+    def _mock_sock(self, responses):
+        """
+        responses: list[bytes],每个是一次 _read_response 调用的完整 RESP 响应。
+        """
+        stream = io.BytesIO(b''.join(responses))
+        sock = MagicMock()
+        sock.makefile.return_value = stream
+        sock.__enter__ = lambda s: s
+        sock.__exit__ = MagicMock(return_value=False)
+        return sock
+
+    def test_execute_with_password_sends_auth(self):
+        """password 非空 → 发送 AUTH 命令。"""
+        client = RedisSocketClient(password='secret', db=0)
+        # AUTH → OK;GET → nil
+        mock_sock = self._mock_sock([b'+OK\r\n', b'$-1\r\n'])
+        with patch('socket.create_connection', return_value=mock_sock):
+            result = client.get('k')
+        self.assertIsNone(result)
+        # sendall 应被调用多次(AUTH + GET)
+        self.assertGreater(mock_sock.sendall.call_count, 0)
+
+    def test_execute_with_db_sends_select(self):
+        """db 非0 → 发送 SELECT 命令。"""
+        client = RedisSocketClient(password='', db=3)
+        # SELECT → OK;GET → value
+        mock_sock = self._mock_sock([b'+OK\r\n', b'$5\r\nvalue\r\n'])
+        with patch('socket.create_connection', return_value=mock_sock):
+            result = client.get('k')
+        self.assertEqual('value', result)
+
+    def test_execute_with_password_and_db(self):
+        """password 非空 且 db 非0 → 发送 AUTH + SELECT。"""
+        client = RedisSocketClient(password='pass', db=2)
+        # AUTH → OK;SELECT → OK;SET → OK
+        mock_sock = self._mock_sock([b'+OK\r\n', b'+OK\r\n', b'+OK\r\n'])
+        with patch('socket.create_connection', return_value=mock_sock):
+            ok = client.set('k', 'v')
+        self.assertTrue(ok)
+
+    def test_execute_without_password_or_db(self):
+        """password 为空、db 为0 → 直接执行命令,不发 AUTH/SELECT。"""
+        client = RedisSocketClient(password='', db=0)
+        mock_sock = self._mock_sock([b'+OK\r\n'])
+        with patch('socket.create_connection', return_value=mock_sock):
+            ok = client.set('k', 'v', ex=300)
+        self.assertTrue(ok)
+
+    def test_delete_returns_integer(self):
+        client = RedisSocketClient()
+        mock_sock = self._mock_sock([b':1\r\n'])
+        with patch('socket.create_connection', return_value=mock_sock):
+            result = client.delete('k')
+        self.assertEqual(1, result)
+
+    def test_set_without_ex(self):
+        client = RedisSocketClient()
+        mock_sock = self._mock_sock([b'+OK\r\n'])
+        with patch('socket.create_connection', return_value=mock_sock):
+            ok = client.set('k', 'v')  # no ex= argument
+        self.assertTrue(ok)
+
+
+if __name__ == '__main__':
+    unittest.main()