import argparse import json import sys import uuid from config import GatewayConfig from constants import DEVICE_INVALID_MESSAGE from mcp_protocol import McpProtocolHandler from public_gateway import PublicGatewayApp from public_server import serve_public from services.api_client import ApiClient from services.auth_client import AuthClient from services.gateway_session_store import GatewaySessionStore from services.scoped_api_client import ScopedApiClient from services.token_store import FileTokenStore, RedisSocketClient, RedisTokenStore from tools.query_order import QueryOrderTool from tools.query_track import QueryTrackTool class GatewayApp: def __init__(self, auth_client=None, api_client=None, token_store=None): self.auth_client = auth_client self.api_client = api_client self.token_store = token_store self._tools = { 'query_order': QueryOrderTool(api_client=api_client), 'query_track': QueryTrackTool(api_client=api_client), } @classmethod def from_config(cls, config, redis_client=None): if config.token_store_type == 'redis': redis = redis_client or RedisSocketClient( host=config.redis_host, port=config.redis_port, db=config.redis_db, password=config.redis_password, timeout=config.timeout_seconds, ) token_store = RedisTokenStore( redis, prefix=config.redis_prefix, session_key=config.session_key, refresh_skew_seconds=config.refresh_skew_seconds, ) elif config.token_store_type == 'file': token_store = FileTokenStore( config.token_store_path, refresh_skew_seconds=config.refresh_skew_seconds, ) else: raise ValueError('unsupported token store type: {0}'.format(config.token_store_type)) auth_client = AuthClient( base_url=config.auth_base_url, client_type=config.client_type, token_store=token_store, timeout=config.timeout_seconds, session_key=config.session_key, ) api_client = ApiClient( base_url=config.tools_base_url, token_store=token_store, timeout=config.timeout_seconds, ) return cls(auth_client=auth_client, api_client=api_client, token_store=token_store) def list_tools(self): return [tool.metadata() for tool in self._tools.values()] def build_request_id(self, request_id=''): request_id = str(request_id or '').strip() if request_id: return request_id return 'rq_{0}'.format(uuid.uuid4().hex[:16]) def ensure_session(self): if self.token_store is None: return session = self.token_store.get() if not session or not session.get('token'): raise RuntimeError(DEVICE_INVALID_MESSAGE) if self.token_store.is_expiring(): if self.auth_client is None: raise RuntimeError('mcp token expiring but auth client missing') self.auth_client.refresh(session['token']) def call_tool(self, name, arguments=None, request_id=''): if name not in self._tools: raise KeyError('tool not registered: {0}'.format(name)) tool = self._tools[name] if getattr(tool, 'requires_session', True): self.ensure_session() arguments = arguments or {} request_id = self.build_request_id(request_id) return tool.call(request_id=request_id, **arguments) def create_protocol_handler(self): return McpProtocolHandler(self) def run_cli(self, argv=None, stdin=None, stdout=None): stdin = stdin or sys.stdin stdout = stdout or sys.stdout parser = argparse.ArgumentParser(prog='mcp-gateway') subparsers = parser.add_subparsers(dest='command', required=True) subparsers.add_parser('list-tools') subparsers.add_parser('serve-stdio') public_parser = subparsers.add_parser('serve-public') public_parser.add_argument('--host', default='0.0.0.0') public_parser.add_argument('--port', type=int, default=8765) call_parser = subparsers.add_parser('call') call_parser.add_argument('--tool', required=True) call_parser.add_argument('--keyword', default='') call_parser.add_argument('--order-id', type=int, default=0) call_parser.add_argument('--order-number', default='') call_parser.add_argument('--tracking-number', default='') call_parser.add_argument('--page', type=int, default=1) call_parser.add_argument('--limit', type=int, default=20) call_parser.add_argument('--request-id', default='') args = parser.parse_args(argv or []) if args.command == 'list-tools': payload = self.list_tools() elif args.command == 'serve-stdio': return self.create_protocol_handler().run_stdio(stdin=stdin, stdout=stdout) elif args.command == 'serve-public': config = GatewayConfig.from_env() redis = RedisSocketClient( host=config.redis_host, port=config.redis_port, db=config.redis_db, password=config.redis_password, timeout=config.timeout_seconds, ) session_store = GatewaySessionStore( redis, prefix=config.redis_prefix, ttl_seconds=config.gateway_session_ttl_seconds, ) public_app = PublicGatewayApp( 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) elif args.command == 'call': tool_args = { 'page': args.page, 'limit': args.limit, } if args.tool == 'query_order': if not args.keyword: raise ValueError('--keyword is required for query_order') tool_args['keyword'] = args.keyword elif args.tool == 'query_track': if args.order_id > 0: tool_args['order_id'] = args.order_id if args.order_number: tool_args['order_number'] = args.order_number if args.tracking_number: tool_args['tracking_number'] = args.tracking_number if args.order_id <= 0 and not args.order_number and not args.tracking_number: raise ValueError('--order-id, --order-number or --tracking-number is required for query_track') else: if args.keyword: tool_args['keyword'] = args.keyword payload = self.call_tool( args.tool, tool_args, request_id=args.request_id, ) else: raise RuntimeError('unsupported command') stdout.write(json.dumps(payload, ensure_ascii=False)) return 0 def main(argv=None): config = GatewayConfig.from_env() app = GatewayApp.from_config(config) return app.run_cli(argv=argv) if __name__ == '__main__': raise SystemExit(main(sys.argv[1:]))